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