Skip to main content

charon_lib/transform/resugar/
reconstruct_vec_boxes.rs

1//! Reconstruct rustc's `vec![..]` lowering based on `Box<MaybeUninit<[T; N]>>`.
2//!
3//! In `inline_selected_functions`, we inline the special `box_assume_init_into_vec_unsafe`
4//! function. After that, a `vec![elems...]` expression ends up looking something like:
5//! ```ignore
6//! let mut box = Box::new_uninit::<[T; N]>();
7//! (((*box).1).0).0 = [elems...];
8//! let box = Box::assume_init(box);
9//! ..
10//! ```
11//! The split between assignment and `assume_init` is for performance. The `assume_init` call is
12//! unsafe, so we rewrite it to use `Box::write` instead, and even `Box::new` if possible.
13//! ```ignore
14//! let box_uninit = Box::new_uninit::<[T; N]>();
15//! let arr = [elems...];
16//! let box = Box::write::<[T; N]>(box_uninit, arr);
17//! ..
18//! ```
19//!
20//! See also: <https://github.com/rust-lang/rust/pull/148190>
21
22use itertools::Itertools;
23use std::collections::HashSet;
24
25use crate::name_matcher::NamePattern;
26use crate::transform::ctx::UllbcPass;
27use crate::transform::{CowBox, TransformCtx};
28use crate::ullbc_ast::*;
29
30pub struct Transform {
31    box_write: Option<FunDeclId>,
32}
33
34struct Rewrite {
35    new_uninit_bid: BlockId,
36    new_uninit_target: BlockId,
37    drop_bid: BlockId,
38    target_bid: BlockId,
39    payload_loc: StmtLoc,
40    move_loc: StmtLoc,
41    arg_move_loc: StmtLoc,
42    span: Span,
43    payload_elems: Vec<Operand>,
44    elem_ty: Ty,
45    len: Box<ConstantExpr>,
46    uninit_box: Place,
47    branched_before_payload: bool,
48    box_array: Place,
49    box_array_generics: GenericArgs,
50    assume_init_generics: GenericArgs,
51    drop_on_unwind: BlockId,
52    assume_init_target: BlockId,
53}
54
55struct PayloadAssign {
56    loc: StmtLoc,
57    span: Span,
58    payload_elems: Vec<Operand>,
59    elem_ty: Ty,
60    len: Box<ConstantExpr>,
61    branched_before_payload: bool,
62}
63
64struct AssumeInitTail {
65    move_loc: StmtLoc,
66    drop_bid: BlockId,
67    target_bid: BlockId,
68    arg_move_loc: StmtLoc,
69    box_array: Place,
70    assume_init_generics: GenericArgs,
71    drop_on_unwind: BlockId,
72    assume_init_target: BlockId,
73}
74
75fn assume_init_fn_ptr<'a>(ctx: &TransformCtx, call: &'a Call) -> Option<&'a FnPtr> {
76    if let FnOperand::Regular(fn_ptr) = &call.func
77        && let FnPtrKind::Fun(FunId::Regular(fid)) = *fn_ptr.kind
78        && ctx.translated.item_name(fid).short_str() == Some("assume_init")
79    {
80        Some(fn_ptr)
81    } else {
82        None
83    }
84}
85
86fn box_inner(ty: &Ty) -> Option<Ty> {
87    let TyKind::Adt(TypeDeclRef {
88        id: TypeId::Builtin(BuiltinTy::Box),
89        generics,
90    }) = ty.kind()
91    else {
92        return None;
93    };
94    Some(generics.types[TypeVarId::from_usize(0)].clone())
95}
96
97fn box_generics(ty: &Ty) -> Option<GenericArgs> {
98    let TyKind::Adt(TypeDeclRef {
99        id: TypeId::Builtin(BuiltinTy::Box),
100        generics,
101    }) = ty.kind()
102    else {
103        return None;
104    };
105    Some((**generics).clone())
106}
107
108/// Given `src`, find the unique statement of the form `src = [elems...]`
109/// where the rvalue is an array aggregate.
110///
111/// Also returns whether a straight-line path from `start` hits a branch before reaching the
112/// assignment. We ignore unwind edges; this matches the later rewrite's ability to erase the
113/// allocation when the normal path to initialization is linear.
114fn find_array_assign(body: &ExprBody, start: BlockId, src_local: LocalId) -> Option<PayloadAssign> {
115    let mut out = None;
116    for (bid, block) in body.body.iter_enumerated() {
117        for (idx, st) in block.statements.iter().enumerate() {
118            let Some((place, Rvalue::Aggregate(AggregateKind::Array(elem_ty, len), elems))) =
119                st.kind.as_assign()
120            else {
121                continue;
122            };
123            if place.local_id() != Some(src_local) {
124                continue;
125            }
126            if out.is_some() {
127                return None;
128            }
129
130            let loc = StmtLoc::new(bid, idx);
131            out = Some(PayloadAssign {
132                loc,
133                span: st.span,
134                payload_elems: elems.clone(),
135                elem_ty: elem_ty.clone(),
136                len: len.clone(),
137                branched_before_payload: branched_before(body, start, loc.block)?,
138            });
139        }
140    }
141    out
142}
143
144fn branched_before(body: &ExprBody, start: BlockId, target: BlockId) -> Option<bool> {
145    let mut block_id = start;
146    let mut visited = HashSet::new();
147    while visited.insert(block_id) {
148        if block_id == target {
149            return Some(false);
150        }
151
152        let block = &body.body[block_id];
153        let targets = block.targets_ignoring_unwind();
154        if targets.len() > 1 {
155            return Some(true);
156        }
157        block_id = targets.into_iter().exactly_one().ok()?;
158    }
159    None
160}
161
162fn unique_target(term: &Terminator) -> Option<BlockId> {
163    term.targets_ignoring_unwind()
164        .into_iter()
165        .exactly_one()
166        .ok()
167}
168
169fn find_next_move_of(
170    body: &ExprBody,
171    mut cursor: StmtLoc,
172    src: &Place,
173) -> Option<(StmtLoc, Place)> {
174    let mut visited = HashSet::new();
175    while visited.insert(cursor) {
176        let block = &body.body[cursor.block];
177        while cursor.statement < block.statements.len() {
178            let st = &body[cursor];
179            if let Some((dst_place, Rvalue::Use(Operand::Move(src_place), _))) = st.kind.as_assign()
180                && src_place == src
181            {
182                return Some((cursor, dst_place.clone()));
183            }
184            cursor = cursor.after();
185        }
186        cursor = StmtLoc::block_start(unique_target(&block.terminator)?);
187    }
188    None
189}
190
191fn find_next_drop(
192    body: &ExprBody,
193    mut block_id: BlockId,
194    dropped_place: &Place,
195) -> Option<(BlockId, BlockId, BlockId)> {
196    let mut visited = HashSet::new();
197    while visited.insert(block_id) {
198        let block = &body.body[block_id];
199        if let TerminatorKind::Drop {
200            place: dropped,
201            target,
202            on_unwind,
203            ..
204        } = &block.terminator.kind
205            && dropped == dropped_place
206        {
207            return Some((block_id, *target, *on_unwind));
208        }
209        block_id = unique_target(&block.terminator)?;
210    }
211    None
212}
213
214fn find_move_in_block(
215    body: &ExprBody,
216    block: BlockId,
217    src: &Place,
218    dst: &Place,
219) -> Option<StmtLoc> {
220    let mut out = None;
221    for (statement, st) in body.body[block].statements.iter().enumerate() {
222        if let Some((dst_place, Rvalue::Use(Operand::Move(src_place), _))) = st.kind.as_assign()
223            && dst_place == dst
224            && src_place == src
225        {
226            if out.is_some() {
227                return None;
228            }
229            out = Some(StmtLoc::new(block, statement));
230        }
231    }
232    out
233}
234
235fn find_assume_init_tail(
236    ctx: &TransformCtx,
237    body: &ExprBody,
238    cursor: StmtLoc,
239    uninit_box: &Place,
240) -> Option<AssumeInitTail> {
241    let (move_loc, moved_box) = find_next_move_of(body, cursor, uninit_box)?;
242    let (drop_bid, target_bid, drop_on_unwind) = find_next_drop(body, move_loc.block, uninit_box)?;
243
244    let target_block = &body.body[target_bid];
245    let (call, assume_init_target, _unwind) = target_block.terminator.kind.as_call()?;
246    let assume_init_fn = assume_init_fn_ptr(ctx, call)?;
247    let [Operand::Move(init_box)] = call.args.as_slice() else {
248        return None;
249    };
250    let arg_move_loc = find_move_in_block(body, target_bid, &moved_box, init_box)?;
251
252    Some(AssumeInitTail {
253        move_loc,
254        drop_bid,
255        target_bid,
256        arg_move_loc,
257        box_array: call.dest.clone(),
258        assume_init_generics: assume_init_fn.generics.as_ref().clone(),
259        drop_on_unwind,
260        assume_init_target: *assume_init_target,
261    })
262}
263
264fn is_new_uninit_call(ctx: &TransformCtx, call: &Call) -> bool {
265    if !call.args.is_empty() {
266        return false;
267    }
268
269    let FnOperand::Regular(fn_ptr) = &call.func else {
270        return false;
271    };
272    let FnPtrKind::Fun(FunId::Regular(fid)) = *fn_ptr.kind else {
273        return false;
274    };
275    ctx.translated.item_name(fid).short_str() == Some("new_uninit")
276}
277
278impl Transform {
279    pub fn new(ctx: &mut TransformCtx) -> CowBox<dyn UllbcPass> {
280        let pat = NamePattern::parse(names::BOX_WRITE_PATTERN).unwrap();
281        let box_write = ctx
282            .translated
283            .item_names
284            .iter()
285            .filter(|(_, name)| pat.matches(&ctx.translated, name))
286            .filter_map(|(id, _)| id.as_fun())
287            .copied()
288            .exactly_one()
289            .ok();
290        CowBox::Owned(Box::new(Transform { box_write }))
291    }
292}
293
294impl UllbcPass for Transform {
295    fn should_run(&self, options: &crate::options::TranslateOptions) -> bool {
296        options.treat_box_as_builtin && !options.monomorphize_with_hax && self.box_write.is_some()
297    }
298
299    fn transform_body(&self, ctx: &mut TransformCtx, body: &mut ExprBody) {
300        // Checked in `should_run`
301        let box_write = self.box_write.unwrap();
302
303        // We are looking for, in (flattened) ULLBC:
304        //
305        // box1 = new_uninit()
306        // ...
307        // ((((*box1)).1).0).0 = [move _4]
308        // box2 = move box1
309        // conditional_drop box1
310        // box3 = move box2
311        // box4 = assume_init(move box3)
312        let rewrites = body
313            .body
314            .iter_enumerated()
315            .filter_map(|(new_uninit_bid, block)| {
316                let TerminatorKind::Call {
317                    call,
318                    target: new_uninit_target,
319                    ..
320                } = &block.terminator.kind
321                else {
322                    return None;
323                };
324
325                if !is_new_uninit_call(ctx, call) {
326                    return None;
327                }
328                // check uninit_box: Box<MaybeUninit<_>>
329                let uninit_box = call.dest.clone();
330                let maybe_uninit_array_ty = box_inner(uninit_box.ty())?;
331                let mu_decl = &ctx
332                    .translated
333                    .type_decls
334                    .get(maybe_uninit_array_ty.as_adt_id()?)?;
335                if mu_decl.item_meta.lang_item.as_ref() != Some(&from_rustc::LangItem::MaybeUninit)
336                {
337                    return None;
338                };
339                let uninit_box_l = uninit_box.local_id()?;
340
341                // (*uninit_box).1.0.0 = [payload_elems...]: [elem_ty; len]
342                let payload = find_array_assign(body, *new_uninit_target, uninit_box_l)?;
343
344                // assume_init(uninit_box2)
345                let tail = find_assume_init_tail(ctx, body, payload.loc.after(), &uninit_box)?;
346                let box_array_generics = box_generics(tail.box_array.ty())?;
347
348                Some(Rewrite {
349                    new_uninit_bid,
350                    new_uninit_target: *new_uninit_target,
351                    drop_bid: tail.drop_bid,
352                    target_bid: tail.target_bid,
353                    payload_loc: payload.loc,
354                    move_loc: tail.move_loc,
355                    arg_move_loc: tail.arg_move_loc,
356                    span: payload.span,
357                    payload_elems: payload.payload_elems,
358                    elem_ty: payload.elem_ty,
359                    len: payload.len,
360                    uninit_box,
361                    branched_before_payload: payload.branched_before_payload,
362                    box_array_generics,
363                    assume_init_generics: tail.assume_init_generics,
364                    drop_on_unwind: tail.drop_on_unwind,
365                    box_array: tail.box_array,
366                    assume_init_target: tail.assume_init_target,
367                })
368            });
369
370        for rw in rewrites.collect::<Vec<_>>() {
371            let array_ty = Ty::mk_array(rw.elem_ty.clone(), *rw.len.clone());
372            let array_local = body.locals.new_var(None, array_ty.clone());
373            let box_array_ty = rw.box_array.ty().clone();
374            let box_array_local = body.locals.new_var(None, box_array_ty.clone());
375
376            let array_lid = array_local.as_local().unwrap();
377            let box_array_lid = box_array_local.as_local().unwrap();
378
379            body[rw.move_loc].kind = StatementKind::Nop;
380            body[rw.arg_move_loc].kind = StatementKind::Nop;
381
382            body.body[rw.payload_loc.block].statements.splice(
383                rw.payload_loc.statement..=rw.payload_loc.statement,
384                [
385                    StatementKind::StorageLive(array_lid),
386                    StatementKind::Assign(
387                        array_local.clone(),
388                        Rvalue::Aggregate(
389                            AggregateKind::Array(rw.elem_ty.clone(), rw.len.clone()),
390                            rw.payload_elems,
391                        ),
392                    ),
393                ]
394                .map(|k| Statement::new(rw.span, k)),
395            );
396
397            let (fn_ptr, args) = if rw.branched_before_payload {
398                (
399                    FnPtr::new(
400                        FnPtrKind::Fun(FunId::Regular(box_write)),
401                        rw.assume_init_generics,
402                    ),
403                    vec![
404                        Operand::Move(rw.uninit_box),
405                        Operand::Move(array_local.clone()),
406                    ],
407                )
408            } else {
409                body.body[rw.new_uninit_bid].terminator.kind = TerminatorKind::Goto {
410                    target: rw.new_uninit_target,
411                };
412                let mut box_new_generics = rw.box_array_generics.clone();
413                box_new_generics.types.pop(); // pop the allocator param
414                (
415                    FnPtr::new(
416                        FnPtrKind::Fun(FunId::Builtin(BuiltinFunId::BoxNew)),
417                        box_new_generics,
418                    ),
419                    vec![Operand::Move(array_local.clone())],
420                )
421            };
422
423            let drop_block = &mut body.body[rw.drop_bid];
424            drop_block.statements.push(Statement::new(
425                rw.span,
426                StatementKind::StorageLive(box_array_lid),
427            ));
428            drop_block.terminator.kind = TerminatorKind::Call {
429                call: Call {
430                    func: FnOperand::Regular(fn_ptr),
431                    args,
432                    dest: box_array_local.clone(),
433                },
434                target: rw.target_bid,
435                on_unwind: rw.drop_on_unwind,
436            };
437
438            let target_block = &mut body.body[rw.target_bid];
439            target_block.statements.push(Statement::new(
440                rw.span,
441                StatementKind::StorageDead(array_lid),
442            ));
443            target_block.statements.push(Statement::new(
444                rw.span,
445                StatementKind::Assign(
446                    rw.box_array,
447                    Rvalue::Use(Operand::Move(box_array_local), WithRetag::No),
448                ),
449            ));
450            target_block.statements.push(Statement::new(
451                rw.span,
452                StatementKind::StorageDead(box_array_lid),
453            ));
454            target_block.terminator.kind = TerminatorKind::Goto {
455                target: rw.assume_init_target,
456            };
457        }
458    }
459}