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_new: Option<FunDeclId>,
32    box_write: Option<FunDeclId>,
33}
34
35struct Rewrite {
36    new_uninit_bid: BlockId,
37    new_uninit_target: BlockId,
38    drop_bid: BlockId,
39    target_bid: BlockId,
40    payload_loc: StmtLoc,
41    move_loc: StmtLoc,
42    arg_move_loc: StmtLoc,
43    span: Span,
44    payload_elems: Vec<Operand>,
45    elem_ty: Ty,
46    elem_ty_is_sized: Option<TraitRef>,
47    len: ConstantExpr,
48    uninit_box: Place,
49    branched_before_payload: bool,
50    box_array: Place,
51    box_new_generics: GenericArgs,
52    assume_init_generics: GenericArgs,
53    drop_on_unwind: BlockId,
54    assume_init_target: BlockId,
55}
56
57struct PayloadAssign {
58    loc: StmtLoc,
59    span: Span,
60    payload_elems: Vec<Operand>,
61    elem_ty: Ty,
62    elem_ty_is_sized: Option<TraitRef>,
63    len: ConstantExpr,
64    branched_before_payload: bool,
65}
66
67struct AssumeInitTail {
68    move_loc: StmtLoc,
69    drop_bid: BlockId,
70    target_bid: BlockId,
71    arg_move_loc: StmtLoc,
72    box_array: Place,
73    assume_init_generics: GenericArgs,
74    drop_on_unwind: BlockId,
75    assume_init_target: BlockId,
76}
77
78fn assume_init_fn_ptr<'a>(ctx: &TransformCtx, call: &'a Call) -> Option<&'a FnPtr> {
79    if let FnOperand::Regular(fn_ptr) = &call.func
80        && let FnPtrKind::Fun(fid) = *fn_ptr.kind
81        && ctx.translated.item_name(fid).short_str() == Some("assume_init")
82    {
83        Some(fn_ptr)
84    } else {
85        None
86    }
87}
88
89fn box_inner(ty: &Ty) -> Option<Ty> {
90    let TypeDeclRef { generics, .. } = ty.as_adt().filter(|tref| tref.is_box())?;
91    Some(generics.types[TypeVarId::from_usize(0)].clone())
92}
93
94/// Given `src`, find the unique statement of the form `src = [elems...]`
95/// where the rvalue is an array aggregate.
96///
97/// Also returns whether a straight-line path from `start` hits a branch before reaching the
98/// assignment. We ignore unwind edges; this matches the later rewrite's ability to erase the
99/// allocation when the normal path to initialization is linear.
100fn find_array_assign(body: &ExprBody, start: BlockId, src_local: LocalId) -> Option<PayloadAssign> {
101    let mut out = None;
102    for (bid, block) in body.body.iter_enumerated() {
103        for (idx, st) in block.statements.iter().enumerate() {
104            let Some((
105                place,
106                Rvalue::Aggregate(AggregateKind::Array(elem_ty, len, elem_ty_is_sized), elems),
107            )) = st.kind.as_assign()
108            else {
109                continue;
110            };
111            if place.local_id() != Some(src_local) {
112                continue;
113            }
114            if out.is_some() {
115                return None;
116            }
117            let loc = StmtLoc::new(bid, idx);
118            out = Some(PayloadAssign {
119                loc,
120                span: st.span,
121                payload_elems: elems.clone(),
122                elem_ty: elem_ty.clone(),
123                elem_ty_is_sized: elem_ty_is_sized.clone(),
124                len: len.clone(),
125                branched_before_payload: branched_before(body, start, loc.block)?,
126            });
127        }
128    }
129    out
130}
131
132fn branched_before(body: &ExprBody, start: BlockId, target: BlockId) -> Option<bool> {
133    let mut block_id = start;
134    let mut visited = HashSet::new();
135    while visited.insert(block_id) {
136        if block_id == target {
137            return Some(false);
138        }
139
140        let block = &body.body[block_id];
141        let targets = block.targets_ignoring_unwind();
142        if targets.len() > 1 {
143            return Some(true);
144        }
145        block_id = targets.into_iter().exactly_one().ok()?;
146    }
147    None
148}
149
150fn unique_target(term: &Terminator) -> Option<BlockId> {
151    term.targets_ignoring_unwind()
152        .into_iter()
153        .exactly_one()
154        .ok()
155}
156
157fn find_next_move_of(
158    body: &ExprBody,
159    mut cursor: StmtLoc,
160    src: &Place,
161) -> Option<(StmtLoc, Place)> {
162    let mut visited = HashSet::new();
163    while visited.insert(cursor) {
164        let block = &body.body[cursor.block];
165        while cursor.statement < block.statements.len() {
166            let st = &body[cursor];
167            if let Some((dst_place, Rvalue::Use(Operand::Move(src_place), _))) = st.kind.as_assign()
168                && src_place == src
169            {
170                return Some((cursor, dst_place.clone()));
171            }
172            cursor = cursor.after();
173        }
174        cursor = StmtLoc::block_start(unique_target(&block.terminator)?);
175    }
176    None
177}
178
179fn find_next_drop(
180    body: &ExprBody,
181    mut block_id: BlockId,
182    dropped_place: &Place,
183) -> Option<(BlockId, BlockId, BlockId)> {
184    let mut visited = HashSet::new();
185    while visited.insert(block_id) {
186        let block = &body.body[block_id];
187        if let TerminatorKind::Drop {
188            place: dropped,
189            target,
190            on_unwind,
191            ..
192        } = &block.terminator.kind
193            && dropped == dropped_place
194        {
195            return Some((block_id, *target, *on_unwind));
196        }
197        block_id = unique_target(&block.terminator)?;
198    }
199    None
200}
201
202fn find_move_in_block(
203    body: &ExprBody,
204    block: BlockId,
205    src: &Place,
206    dst: &Place,
207) -> Option<StmtLoc> {
208    let mut out = None;
209    for (statement, st) in body.body[block].statements.iter().enumerate() {
210        if let Some((dst_place, Rvalue::Use(Operand::Move(src_place), _))) = st.kind.as_assign()
211            && dst_place == dst
212            && src_place == src
213        {
214            if out.is_some() {
215                return None;
216            }
217            out = Some(StmtLoc::new(block, statement));
218        }
219    }
220    out
221}
222
223fn find_assume_init_tail(
224    ctx: &TransformCtx,
225    body: &ExprBody,
226    cursor: StmtLoc,
227    uninit_box: &Place,
228) -> Option<AssumeInitTail> {
229    let (move_loc, moved_box) = find_next_move_of(body, cursor, uninit_box)?;
230    let (drop_bid, target_bid, drop_on_unwind) = find_next_drop(body, move_loc.block, uninit_box)?;
231
232    let target_block = &body.body[target_bid];
233    let (call, assume_init_target, _unwind) = target_block.terminator.kind.as_call()?;
234    let assume_init_fn = assume_init_fn_ptr(ctx, call)?;
235    let [Operand::Move(init_box)] = call.args.as_slice() else {
236        return None;
237    };
238    let arg_move_loc = find_move_in_block(body, target_bid, &moved_box, init_box)?;
239
240    Some(AssumeInitTail {
241        move_loc,
242        drop_bid,
243        target_bid,
244        arg_move_loc,
245        box_array: call.dest.clone(),
246        assume_init_generics: assume_init_fn.generics.as_ref().clone(),
247        drop_on_unwind,
248        assume_init_target: *assume_init_target,
249    })
250}
251
252fn is_new_uninit_call(ctx: &TransformCtx, call: &Call) -> bool {
253    if !call.args.is_empty() {
254        return false;
255    }
256
257    let FnOperand::Regular(fn_ptr) = &call.func else {
258        return false;
259    };
260    let FnPtrKind::Fun(fid) = *fn_ptr.kind else {
261        return false;
262    };
263    ctx.translated.item_name(fid).short_str() == Some("new_uninit")
264}
265
266impl Transform {
267    pub fn new(ctx: &mut TransformCtx) -> CowBox<dyn UllbcPass> {
268        let box_new = ctx
269            .translated
270            .fun_decls
271            .iter()
272            .filter(|decl| decl.item_meta.diagnostic_item.as_deref() == Some("box_new"))
273            .map(|decl| decl.def_id)
274            .exactly_one()
275            .ok();
276        let pat = NamePattern::parse(names::BOX_WRITE_PATTERN).unwrap();
277        let box_write = ctx
278            .translated
279            .item_names
280            .iter()
281            .filter(|(_, name)| pat.matches(&ctx.translated, name))
282            .filter_map(|(id, _)| id.as_fun())
283            .copied()
284            .exactly_one()
285            .ok();
286        CowBox::Owned(Box::new(Transform { box_new, box_write }))
287    }
288}
289
290impl UllbcPass for Transform {
291    fn should_run(&self, options: &crate::options::TranslateOptions) -> bool {
292        options.treat_box_as_builtin
293            && !options.monomorphize_with_hax
294            && self.box_new.is_some()
295            && self.box_write.is_some()
296    }
297
298    fn transform_body(&self, ctx: &mut TransformCtx, body: &mut ExprBody) {
299        // Checked in `should_run`
300        let box_new = self.box_new.unwrap();
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 maybe_uninit_ref = maybe_uninit_array_ty.as_adt()?;
332                let mu_decl = &ctx.translated.type_decls.get(maybe_uninit_ref.id)?;
333                if mu_decl.item_meta.lang_item.as_ref() != Some(&from_rustc::LangItem::MaybeUninit)
334                {
335                    return None;
336                };
337                // `MaybeUninit<T>` and `Box::new<T>` have the same generic parameters (`T: Sized`).
338                let box_new_generics = maybe_uninit_ref.generics.as_ref().clone();
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                Some(Rewrite {
347                    new_uninit_bid,
348                    new_uninit_target: *new_uninit_target,
349                    drop_bid: tail.drop_bid,
350                    target_bid: tail.target_bid,
351                    payload_loc: payload.loc,
352                    move_loc: tail.move_loc,
353                    arg_move_loc: tail.arg_move_loc,
354                    span: payload.span,
355                    payload_elems: payload.payload_elems,
356                    elem_ty: payload.elem_ty,
357                    elem_ty_is_sized: payload.elem_ty_is_sized,
358                    len: payload.len,
359                    uninit_box,
360                    branched_before_payload: payload.branched_before_payload,
361                    box_new_generics,
362                    assume_init_generics: tail.assume_init_generics,
363                    drop_on_unwind: tail.drop_on_unwind,
364                    box_array: tail.box_array,
365                    assume_init_target: tail.assume_init_target,
366                })
367            });
368
369        for rw in rewrites.collect::<Vec<_>>() {
370            let array_ty = Ty::mk_array(
371                rw.elem_ty.clone(),
372                rw.len.clone(),
373                rw.elem_ty_is_sized.clone(),
374            );
375            let array_local = body.locals.new_var(None, array_ty.clone());
376            let box_array_ty = rw.box_array.ty().clone();
377            let box_array_local = body.locals.new_var(None, box_array_ty.clone());
378
379            let array_lid = array_local.as_local().unwrap();
380            let box_array_lid = box_array_local.as_local().unwrap();
381
382            body[rw.move_loc].kind = StatementKind::Nop;
383            body[rw.arg_move_loc].kind = StatementKind::Nop;
384
385            body.body[rw.payload_loc.block].statements.splice(
386                rw.payload_loc.statement..=rw.payload_loc.statement,
387                [
388                    StatementKind::StorageLive(array_lid),
389                    StatementKind::Assign(
390                        array_local.clone(),
391                        Rvalue::Aggregate(
392                            AggregateKind::Array(
393                                rw.elem_ty.clone(),
394                                rw.len.clone(),
395                                rw.elem_ty_is_sized,
396                            ),
397                            rw.payload_elems,
398                        ),
399                    ),
400                ]
401                .map(|k| Statement::new(rw.span, k)),
402            );
403
404            let (fn_ptr, args) = if rw.branched_before_payload {
405                (
406                    FnPtr::new(FnPtrKind::Fun(box_write), rw.assume_init_generics),
407                    vec![
408                        Operand::Move(rw.uninit_box),
409                        Operand::Move(array_local.clone()),
410                    ],
411                )
412            } else {
413                body.body[rw.new_uninit_bid].terminator.kind = TerminatorKind::Goto {
414                    target: rw.new_uninit_target,
415                };
416                (
417                    FnPtr::new(FnPtrKind::Fun(box_new), rw.box_new_generics),
418                    vec![Operand::Move(array_local.clone())],
419                )
420            };
421
422            let drop_block = &mut body.body[rw.drop_bid];
423            drop_block.statements.push(Statement::new(
424                rw.span,
425                StatementKind::StorageLive(box_array_lid),
426            ));
427            drop_block.terminator.kind = TerminatorKind::Call {
428                call: Call {
429                    func: FnOperand::Regular(fn_ptr),
430                    args,
431                    dest: box_array_local.clone(),
432                },
433                target: rw.target_bid,
434                on_unwind: rw.drop_on_unwind,
435            };
436
437            let target_block = &mut body.body[rw.target_bid];
438            target_block.statements.push(Statement::new(
439                rw.span,
440                StatementKind::StorageDead(array_lid),
441            ));
442            target_block.statements.push(Statement::new(
443                rw.span,
444                StatementKind::Assign(
445                    rw.box_array,
446                    Rvalue::Use(Operand::Move(box_array_local), WithRetag::No),
447                ),
448            ));
449            target_block.statements.push(Statement::new(
450                rw.span,
451                StatementKind::StorageDead(box_array_lid),
452            ));
453            target_block.terminator.kind = TerminatorKind::Goto {
454                target: rw.assume_init_target,
455            };
456        }
457    }
458}