Skip to main content

rustc_middle/ty/
typetree.rs

1use rustc_ast::expand::typetree::{FncTree, Kind, Type, TypeTree};
2use rustc_span::bug;
3use tracing::trace;
4
5use crate::ty::consts::ConstExt;
6use crate::ty::context::TyCtxt;
7use crate::ty::{self, Ty};
8
9/// Generate TypeTree information for autodiff.
10/// This function creates TypeTree metadata that describes the memory layout
11/// of function parameters and return types for Enzyme autodiff.
12pub fn fnc_typetrees<'tcx>(tcx: TyCtxt<'tcx>, fn_ty: Ty<'tcx>) -> FncTree {
13    // Check if TypeTrees are disabled via NoTT flag
14    if tcx.sess.opts.unstable_opts.autodiff.contains(&rustc_session::config::AutoDiff::NoTT) {
15        return FncTree { args: ::alloc::vec::Vec::new()vec![], ret: TypeTree::new() };
16    }
17
18    // Check if this is actually a function type
19    if !fn_ty.is_fn() {
20        return FncTree { args: ::alloc::vec::Vec::new()vec![], ret: TypeTree::new() };
21    }
22
23    // Get the function signature
24    let fn_sig = fn_ty.fn_sig(tcx);
25    let sig = tcx.instantiate_bound_regions_with_erased(fn_sig);
26
27    // Create TypeTrees for each input parameter
28    let mut args = ::alloc::vec::Vec::new()vec![];
29    for ty in sig.inputs().iter() {
30        let type_tree = typetree_from_ty(tcx, *ty);
31        args.push(type_tree);
32    }
33
34    // Create TypeTree for return type
35    let ret = typetree_from_ty(tcx, sig.output());
36
37    let f = FncTree { args, ret };
38    f
39}
40
41/// Generate a TypeTree for a specific type.
42/// Mainly a convenience wrapper around the actual implementation.
43pub fn typetree_from_ty<'tcx>(tcx: TyCtxt<'tcx>, ty: Ty<'tcx>) -> TypeTree {
44    if !tcx.sess.opts.unstable_opts.autodiff.contains(&rustc_session::config::AutoDiff::Enable) {
45        return TypeTree::new();
46    }
47    if tcx.sess.opts.unstable_opts.autodiff.contains(&rustc_session::config::AutoDiff::NoTT) {
48        return TypeTree::new();
49    }
50    let mut visited = Vec::new();
51    typetree_from_ty_impl_inner(tcx, ty, 0, &mut visited, false)
52}
53
54/// Maximum recursion depth for TypeTree generation to prevent stack overflow
55/// from pathological deeply nested types. Combined with cycle detection.
56const MAX_TYPETREE_DEPTH: usize = 6;
57
58fn handle_indirection<'a>(
59    ty: Ty<'a>,
60    tcx: TyCtxt<'a>,
61    depth: usize,
62    visited: &mut Vec<Ty<'a>>,
63) -> TypeTree {
64    let Some(inner_ty) = ty.builtin_deref(true) else {
65        ::rustc_span::macros::bug_impl(None,
    format_args!("incorrect autodiff typetree handling for type: {0}", ty),
    Location::caller());bug!("incorrect autodiff typetree handling for type: {}", ty);
66    };
67    // A pointer to a slice-tailed DST is a fat pointer `{data, len}`. `RustSlice` describes both
68    // LLVM arguments, while its child describes the memory reached through `data`.
69    let typing_env = ty::TypingEnv::fully_monomorphized();
70    if let ty::Slice(element_ty) = tcx.struct_tail_for_codegen(inner_ty, typing_env).kind() {
71        // `layout.size` here is the sized prefix of `inner_ty`, not the slice element size.
72        // Direct slices, transparent wrappers (`OsStr`), and ZST-prefixed DSTs have no byte
73        // offset to preserve. Nonzero prefixes (e.g. `Header<[f32]>`) keep field offsets.
74        // ZST elements still take this path and yield an empty child TypeTree (size 0).
75        let child = if tcx
76            .layout_of(typing_env.as_query_input(inner_ty))
77            .is_ok_and(|layout| layout.size.bytes() == 0)
78        {
79            typetree_from_ty_impl_inner(tcx, *element_ty, depth + 1, visited, false)
80        } else {
81            typetree_from_ty_impl_inner(tcx, inner_ty, depth + 1, visited, true)
82        };
83        return TypeTree(::alloc::boxed::box_assume_init_into_vec_unsafe(::alloc::intrinsics::write_box_via_move(::alloc::boxed::Box::new_uninit(),
        [Type {
                    offset: -1,
                    size: tcx.data_layout.pointer_size().bytes_usize(),
                    kind: Kind::RustSlice,
                    child,
                }]))vec![Type {
84            offset: -1,
85            size: tcx.data_layout.pointer_size().bytes_usize(),
86            kind: Kind::RustSlice,
87            child,
88        }]);
89    }
90
91    let child = typetree_from_ty_impl_inner(tcx, inner_ty, depth + 1, visited, true);
92    return TypeTree(::alloc::boxed::box_assume_init_into_vec_unsafe(::alloc::intrinsics::write_box_via_move(::alloc::boxed::Box::new_uninit(),
        [Type {
                    offset: -1,
                    size: tcx.data_layout.pointer_size().bytes_usize(),
                    kind: Kind::Pointer,
                    child,
                }]))vec![Type {
93        offset: -1,
94        size: tcx.data_layout.pointer_size().bytes_usize(),
95        kind: Kind::Pointer,
96        child,
97    }]);
98}
99
100/// Internal implementation with context about whether this is for a reference target.
101fn typetree_from_ty_impl_inner<'tcx>(
102    tcx: TyCtxt<'tcx>,
103    ty: Ty<'tcx>,
104    depth: usize,
105    visited: &mut Vec<Ty<'tcx>>,
106    is_reference_target: bool,
107) -> TypeTree {
108    if depth >= MAX_TYPETREE_DEPTH {
109        {
    use ::tracing::__macro_support::Callsite as _;
    static __CALLSITE: ::tracing::callsite::DefaultCallsite =
        {
            static META: ::tracing::Metadata<'static> =
                {
                    ::tracing_core::metadata::Metadata::new("event /rustc-dev/d080e7dff1b0fc54541545252818f8cccf995d05/compiler/rustc_middle/src/ty/typetree.rs:109",
                        "rustc_middle::ty::typetree", ::tracing::Level::TRACE,
                        ::tracing_core::__macro_support::Option::Some("/rustc-dev/d080e7dff1b0fc54541545252818f8cccf995d05/compiler/rustc_middle/src/ty/typetree.rs"),
                        ::tracing_core::__macro_support::Option::Some(109u32),
                        ::tracing_core::__macro_support::Option::Some("rustc_middle::ty::typetree"),
                        ::tracing_core::field::FieldSet::new(&["message"],
                            ::tracing_core::callsite::Identifier(&__CALLSITE)),
                        ::tracing::metadata::Kind::EVENT)
                };
            ::tracing::callsite::DefaultCallsite::new(&META)
        };
    let enabled =
        ::tracing::Level::TRACE <= ::tracing::level_filters::STATIC_MAX_LEVEL
                &&
                ::tracing::Level::TRACE <=
                    ::tracing::level_filters::LevelFilter::current() &&
            {
                let interest = __CALLSITE.interest();
                !interest.is_never() &&
                    ::tracing::__macro_support::__is_enabled(__CALLSITE.metadata(),
                        interest)
            };
    if enabled {
        (|value_set: ::tracing::field::ValueSet|
                    {
                        let meta = __CALLSITE.metadata();
                        ::tracing::Event::dispatch(meta, &value_set);
                        ;
                    })({
                #[allow(unused_imports)]
                use ::tracing::field::{debug, display, Value};
                __CALLSITE.metadata().fields().value_set_all(&[(::tracing::__macro_support::Option::Some(&format_args!("typetree depth limit {0} reached for type: {1}",
                                                    MAX_TYPETREE_DEPTH, ty) as &dyn ::tracing::field::Value))])
            });
    } else { ; }
};trace!("typetree depth limit {} reached for type: {}", MAX_TYPETREE_DEPTH, ty);
110        return TypeTree::new();
111    }
112
113    if visited.contains(&ty) {
114        return TypeTree::new();
115    }
116    visited.push(ty);
117
118    let tree = match ty.kind() {
119        // Direct slices are handled by `handle_indirection`. This arm describes a slice tail while
120        // recursing through a prefixed DST, so its caller can add the field offset.
121        ty::Slice(element_ty) => {
122            typetree_from_ty_impl_inner(tcx, *element_ty, depth + 1, visited, false)
123        }
124        ty::Ref(..) | ty::RawPtr(..) => handle_indirection(ty, tcx, depth, visited),
125        ty::Adt(def, _) if def.is_box() => handle_indirection(ty, tcx, depth, visited),
126        ty::Array(element_ty, len_const) => {
127            let len = len_const.try_to_target_usize(tcx).unwrap_or(0);
128            if len == 0 {
129                TypeTree::new()
130            } else {
131                let element_tree =
132                    typetree_from_ty_impl_inner(tcx, *element_ty, depth + 1, visited, false);
133                let mut types = Vec::new();
134                for elem_type in &element_tree.0 {
135                    types.push(Type::from_ty(-1, elem_type));
136                }
137
138                TypeTree(types)
139            }
140        }
141        ty::Tuple(tuple_types) => {
142            if tuple_types.is_empty() {
143                TypeTree::new()
144            } else {
145                let mut types = Vec::new();
146                let mut current_offset = 0;
147
148                for tuple_ty in tuple_types.iter() {
149                    let element_tree =
150                        typetree_from_ty_impl_inner(tcx, tuple_ty, depth + 1, visited, false);
151
152                    let element_layout = tcx
153                        .layout_of(ty::TypingEnv::fully_monomorphized().as_query_input(tuple_ty))
154                        .ok()
155                        .map(|layout| layout.size.bytes_usize())
156                        .unwrap_or(0);
157
158                    for elem_type in &element_tree.0 {
159                        let offset = if elem_type.offset == -1 {
160                            current_offset as isize
161                        } else {
162                            current_offset as isize + elem_type.offset
163                        };
164                        types.push(Type::from_ty(offset, elem_type));
165                    }
166
167                    current_offset += element_layout;
168                }
169
170                TypeTree(types)
171            }
172        }
173        ty::Adt(adt_def, args) if adt_def.is_struct() => {
174            let struct_layout =
175                tcx.layout_of(ty::TypingEnv::fully_monomorphized().as_query_input(ty));
176            if let Ok(layout) = struct_layout {
177                let mut types = Vec::new();
178
179                for (field_idx, field_def) in adt_def.all_fields().enumerate() {
180                    let field_ty = field_def.ty(tcx, args);
181                    let field_tree = typetree_from_ty_impl_inner(
182                        tcx,
183                        field_ty.skip_norm_wip(),
184                        depth + 1,
185                        visited,
186                        false,
187                    );
188
189                    let field_offset = layout.fields.offset(field_idx).bytes_usize();
190
191                    for elem_type in &field_tree.0 {
192                        let offset = if elem_type.offset == -1 {
193                            field_offset as isize
194                        } else {
195                            field_offset as isize + elem_type.offset
196                        };
197                        types.push(Type::from_ty(offset, elem_type));
198                    }
199                }
200
201                TypeTree(types)
202            } else {
203                TypeTree::new()
204            }
205        }
206        ty::Char | ty::Bool | ty::Infer(ty::IntVar(_)) | ty::Int(_) | ty::Uint(_) => {
207            let kind = Kind::Integer;
208            let size = ty.primitive_size(tcx).bytes_usize();
209            let offset = if is_reference_target { 0 } else { -1 };
210            TypeTree(::alloc::boxed::box_assume_init_into_vec_unsafe(::alloc::intrinsics::write_box_via_move(::alloc::boxed::Box::new_uninit(),
        [Type { offset, size, kind, child: TypeTree::new() }]))vec![Type { offset, size, kind, child: TypeTree::new() }])
211        }
212        ty::Float(_) | ty::Infer(ty::FloatVar(_)) => {
213            let (enzyme_ty, size) = match ty {
214                x if x == tcx.types.f16 => (Kind::Half, 2),
215                x if x == tcx.types.f32 => (Kind::Float, 4),
216                x if x == tcx.types.f64 => (Kind::Double, 8),
217                x if x == tcx.types.f128 => (Kind::F128, 16),
218                _ => ::rustc_span::macros::bug_impl(None,
    format_args!("Unexpected floating point type: {0:?}", ty),
    Location::caller())bug!("Unexpected floating point type: {:?}", ty),
219            };
220            let offset = if is_reference_target { 0 } else { -1 };
221            TypeTree(::alloc::boxed::box_assume_init_into_vec_unsafe(::alloc::intrinsics::write_box_via_move(::alloc::boxed::Box::new_uninit(),
        [Type { offset, size, kind: enzyme_ty, child: TypeTree::new() }]))vec![Type { offset, size, kind: enzyme_ty, child: TypeTree::new() }])
222        }
223        _ => TypeTree::new(),
224    };
225
226    let popped = visited.pop();
227    if true {
    {
        match (&popped, &Some(ty)) {
            (left_val, right_val) => {
                if !(*left_val == *right_val) {
                    let kind = ::core::panicking::AssertKind::Eq;
                    ::core::panicking::assert_failed(kind, &*left_val,
                        &*right_val, ::core::option::Option::None);
                }
            }
        }
    };
};debug_assert_eq!(popped, Some(ty));
228    tree
229}