Skip to main content

charon_lib/transform/normalize/
transform_dyn_trait_calls.rs

1//! Transform method calls on `&dyn Trait` to vtable function pointer calls.
2//!
3//! This pass converts direct method calls on trait objects into calls through vtable
4//! function pointers. For example:
5//!
6//! ```rust,ignore
7//! let x: &dyn Trait = &obj;
8//! x.method(args);
9//! ```
10//!
11//! is transformed from:
12//! ```text
13//! @0 := call <dyn Trait as Trait>::method(x, args)
14//! ```
15//! to:
16//! ```text
17//! @0 := (move (*@receiver.ptr_metadata).method_check)(move (@receiver), move (@args)) // Call through function pointer
18//! ```
19use super::super::ctx::UllbcPass;
20use crate::{
21    errors::Error,
22    formatter::IntoFormatter,
23    pretty::FmtWithCtx,
24    raise_error, register_error,
25    transform::{
26        TransformCtx,
27        ctx::{BodyTransformCtx, UllbcStatementTransformCtx},
28    },
29    ullbc_ast::*,
30};
31
32/// Transform a call to a trait method on a dyn trait object
33fn transform_dyn_trait_call(
34    ctx: &mut UllbcStatementTransformCtx<'_>,
35    call: &mut Call,
36) -> Result<(), Error> {
37    let fmt_ctx = &ctx.ctx.into_fmt();
38
39    // Detect if this call should be transformed
40    let FnOperand::Regular(fn_ptr) = &call.func else {
41        return Ok(()); // Not a regular function call
42    };
43    let FnPtrKind::Trait(trait_ref, method_id) = fn_ptr.kind.as_ref() else {
44        return Ok(()); // Not a trait method call
45    };
46    let mut dyn_proof = trait_ref;
47    let mut supertrait_path = vec![];
48    while let TraitRefKind::ParentClause(parent, clause_id) = &dyn_proof.kind {
49        supertrait_path.push(*clause_id);
50        dyn_proof = parent;
51    }
52    let TraitRefKind::Dyn = &dyn_proof.kind else {
53        return Ok(()); // Not a dyn trait trait call
54    };
55    supertrait_path.reverse();
56
57    // The drop glue of the `Destruct` implied clause is the drop function, which
58    // we must manually desugar, because `Destruct` is not dyn-compatible.
59    let is_drop_glue = ctx
60        .ctx
61        .translated
62        .trait_decls
63        .get(trait_ref.trait_id())
64        .and_then(|decl| decl.item_meta.lang_item.as_ref())
65        == Some(&from_rustc::LangItem::Destruct);
66
67    let (vtable_tref, target_field) = if is_drop_glue {
68        supertrait_path.clear(); // we actually don't need to go up; the drop is here!
69        (dyn_proof, VTableField::Drop)
70    } else {
71        (trait_ref, VTableField::Method(*method_id))
72    };
73
74    // Get the type of the vtable struct.
75    let vtable_decl_ref: TypeDeclRef = {
76        // Get the trait declaration by its ID
77        let Some(trait_decl) = ctx.ctx.translated.trait_decls.get(vtable_tref.trait_id()) else {
78            return Ok(()); // Unknown trait
79        };
80        // Get vtable ref from definition for correct ID.
81        let Some(vtable_ty) = &trait_decl.vtable else {
82            raise_error!(
83                ctx.ctx,
84                ctx.span,
85                "Found a `dyn Trait` method call for non-dyn-compatible trait `{}`!",
86                vtable_tref.trait_id().with_ctx(fmt_ctx)
87            );
88        };
89        vtable_ty.clone().substitute_with_tref(vtable_tref)
90    };
91
92    let Some(vtable_decl) = ctx.ctx.translated.type_decls.get(vtable_decl_ref.id) else {
93        return Ok(()); // Missing data
94    };
95
96    let TypeDeclKind::Struct(fields) = &vtable_decl.kind else {
97        return Ok(()); // Missing data
98    };
99    let TypeSource::VTable { field_map, .. } = &vtable_decl.src else {
100        return Ok(()); // Weird
101    };
102    // Retrieve the target field from the vtable struct definition.
103    let Some((method_field_id, _)) = field_map
104        .iter_enumerated()
105        .find(|(_, field)| **field == target_field)
106    else {
107        let vtable_name = vtable_decl_ref.id.with_ctx(fmt_ctx).to_string();
108        raise_error!(
109            ctx.ctx,
110            ctx.span,
111            "Could not determine method index for method {} in vtable {}",
112            method_id,
113            vtable_name
114        );
115    };
116
117    let method_field = &fields[method_field_id];
118    let method_ty = method_field
119        .ty
120        .clone()
121        .substitute(&vtable_decl_ref.generics);
122
123    // Get the receiver (first argument).
124    if call.args.is_empty() {
125        raise_error!(ctx.ctx, ctx.span, "Dyn trait call has no arguments!");
126    }
127    let mut dyn_trait_place = match &call.args[0] {
128        Operand::Copy(place) | Operand::Move(place) => place.clone(),
129        Operand::Const(_) => {
130            panic!("Unexpected constant as receiver for dyn trait method call")
131        }
132    };
133
134    // Rustc may move the unsized first argument of by-value dyn call into a temporary local,
135    // e.g. for `StructWithTail<dyn T>>` and `b.field.by_value()`, we get `_15 = move (*b).field; by_value(move _15)`.
136    // That's an unsized local, which we is not supposed to happen. To avoid that we inline the
137    // place back into the call, giving us: `by_value(move (*b).field)`.
138    if let TyKind::DynTrait(..) = dyn_trait_place.ty().kind()
139        && let PlaceKind::Local(local) = dyn_trait_place.kind
140        && let Some(last) = ctx.statements.last()
141        && let StatementKind::Assign(dest, Rvalue::Use(src, _)) = &last.kind
142        && let Operand::Copy(place) | Operand::Move(place) = src
143        && dest.local_id() == Some(local)
144    {
145        call.args[0] = src.clone();
146        dyn_trait_place = place.clone();
147        ctx.statements.pop();
148    }
149
150    // `dyn Trait` has a carveout for unsized `self` types. It works by calling the method on a place
151    // of type `dyn Trait` directly. The actual shim we generate expects `*mut dyn Trait`, so we must
152    // borrow that unsized place. See `dyn/mono-dyn-call-by-value.out`.
153    if let TyKind::DynTrait(..) = dyn_trait_place.ty().kind() {
154        let ptr_ty = TyKind::RawPtr(dyn_trait_place.ty().clone(), RefKind::Mut).into_ty();
155        let rvalue = ctx.raw_borrow(dyn_trait_place, RefKind::Mut);
156        dyn_trait_place = ctx.fresh_var(None, ptr_ty);
157        ctx.insert_assn_stmt(dyn_trait_place.clone(), rvalue);
158        call.args[0] = Operand::Move(dyn_trait_place.clone());
159    }
160
161    let dyn_pred = dyn_proof.trait_decl_ref.clone().erase();
162    let dyn_ty = &dyn_pred.generics.types[0];
163    let PtrMetadata::VTable(receiver_vtable_ref) = dyn_ty.get_ptr_metadata(&ctx.ctx.translated)
164    else {
165        raise_error!(
166            ctx.ctx,
167            ctx.span,
168            "Dyn trait receiver does not have vtable metadata"
169        );
170    };
171
172    let receiver_vtable_ty = TyKind::Adt(receiver_vtable_ref).into_ty();
173    let ptr_to_vtable_ty = Ty::new(TyKind::RawPtr(receiver_vtable_ty.clone(), RefKind::Shared));
174    let mut method_vtable_place = dyn_trait_place
175        .project(ProjectionElem::PtrMetadata, ptr_to_vtable_ty)
176        .project(ProjectionElem::Deref, receiver_vtable_ty);
177
178    for clause_id in supertrait_path {
179        let current_vtable_ref = method_vtable_place
180            .ty()
181            .as_adt()
182            .expect("vtable place should have an ADT type");
183        let current_vtable = ctx
184            .ctx
185            .translated
186            .type_decls
187            .get(current_vtable_ref.id)
188            .expect("vtable declaration should have been translated");
189        let TypeDeclKind::Struct(fields) = &current_vtable.kind else {
190            panic!("vtable declaration should be a struct")
191        };
192        let TypeSource::VTable { supertrait_map, .. } = &current_vtable.src else {
193            panic!("vtable declaration should have a vtable source")
194        };
195        let Some(supertrait_field_id) = supertrait_map[clause_id] else {
196            raise_error!(
197                ctx.ctx,
198                ctx.span,
199                "Dyn trait proof uses a parent clause without a vtable"
200            );
201        };
202        let field_ty = fields[supertrait_field_id]
203            .ty
204            .clone()
205            .substitute(&current_vtable_ref.generics);
206        method_vtable_place = method_vtable_place
207            .project(ProjectionElem::Field(None, supertrait_field_id), field_ty)
208            .deref();
209    }
210    let method_field_place =
211        method_vtable_place.project(ProjectionElem::Field(None, method_field_id), method_ty);
212
213    let fn_ptr_place = if ctx.ctx.options.monomorphize_with_hax {
214        // In mono mode, the vtable contains erased function pointers, cast to `*const ()`.
215        // This casts back to the expected signature.
216        let real_sig_ty = TyKind::FnPtr(RegionBinder::empty(FunSig {
217            is_unsafe: true,
218            abi: Abi::rust(),
219            is_variadic: false,
220            inputs: call.args.iter().map(|op| op.ty().clone()).collect(),
221            output: call.dest.ty.clone(),
222        }))
223        .into_ty();
224        let fn_ptr_place = ctx.fresh_var(None, real_sig_ty);
225        let rval_cast = Rvalue::UnaryOp(
226            UnOp::Cast(CastKind::RawPtr(
227                method_field_place.ty().clone(),
228                fn_ptr_place.ty().clone(),
229            )),
230            Operand::Copy(method_field_place),
231        );
232        ctx.insert_assn_stmt(fn_ptr_place.clone(), rval_cast);
233        fn_ptr_place
234    } else {
235        method_field_place
236    };
237
238    // Transform the original call to use the function pointer
239    call.func = FnOperand::Dynamic(Operand::Copy(fn_ptr_place));
240
241    Ok(())
242}
243
244pub struct Transform;
245impl UllbcPass for Transform {
246    fn transform_function(&self, ctx: &mut TransformCtx, decl: &mut FunDecl) {
247        decl.transform_ullbc_terminators(ctx, |ctx, term| {
248            if let TerminatorKind::Call { call, .. } = &mut term.kind {
249                let _ = transform_dyn_trait_call(ctx, call);
250            }
251        });
252    }
253}