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
9pub fn fnc_typetrees<'tcx>(tcx: TyCtxt<'tcx>, fn_ty: Ty<'tcx>) -> FncTree {
13 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 if !fn_ty.is_fn() {
20 return FncTree { args: ::alloc::vec::Vec::new()vec![], ret: TypeTree::new() };
21 }
22
23 let fn_sig = fn_ty.fn_sig(tcx);
25 let sig = tcx.instantiate_bound_regions_with_erased(fn_sig);
26
27 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 let ret = typetree_from_ty(tcx, sig.output());
36
37 let f = FncTree { args, ret };
38 f
39}
40
41pub 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
54const 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 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 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
100fn 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 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}