1use rustc_ast::token::{Delimiter, Token, TokenKind};
2use rustc_ast::tokenstream::{DelimSpan, Spacing, TokenStream, TokenTree};
3use rustc_ast::{AttrItem, ast};
4use rustc_expand::base::{Annotatable, ExtCtxt};
5use rustc_session::config::Offload;
6use rustc_span::{Ident, Span, sym};
7use thin_vec::thin_vec;
8
9use crate::diagnostics;
10
11fn compile_for_device(ecx: &mut ExtCtxt<'_>) -> bool {
12 ecx.sess.opts.unstable_opts.offload.contains(&Offload::Device)
13}
14
15fn outer_normal_attr(
16 kind: &Box<rustc_ast::NormalAttr>,
17 id: rustc_ast::AttrId,
18 span: Span,
19) -> rustc_ast::Attribute {
20 let style = rustc_ast::AttrStyle::Outer;
21 let kind = rustc_ast::AttrKind::Normal(kind.clone());
22 rustc_ast::Attribute { kind, id, style, span }
23}
24
25fn extract_fn(
26 item: &Annotatable,
27) -> Option<(ast::Visibility, ast::FnSig, Ident, ast::Generics, Option<Box<ast::Block>>)> {
28 match item {
29 Annotatable::Item(iitem) => match &iitem.kind {
30 ast::ItemKind::Fn(ast::Fn { sig, ident, generics, body, .. }) => {
31 Some((iitem.vis.clone(), sig.clone(), *ident, generics.clone(), body.clone()))
32 }
33 _ => None,
34 },
35 _ => None,
36 }
37}
38
39pub(crate) fn expand_kernel(
69 ecx: &mut ExtCtxt<'_>,
70 expand_span: Span,
71 _meta_item: &ast::MetaItem,
72 item: Annotatable,
73) -> Vec<Annotatable> {
74 let dcx = ecx.sess.dcx();
75
76 let Some((vis, sig, ident, generics, body)) = extract_fn(&item) else {
77 dcx.emit_err(diagnostics::AutoDiffInvalidApplication { span: item.span() });
78 return ::alloc::boxed::box_assume_init_into_vec_unsafe(::alloc::intrinsics::write_box_via_move(::alloc::boxed::Box::new_uninit(),
[item]))vec![item];
79 };
80
81 let span = ecx.with_def_site_ctxt(expand_span);
82
83 let mut device_fn = Box::new(ast::Fn {
85 defaultness: ast::Defaultness::Implicit,
86 sig: sig.clone(),
87 ident,
88 generics: generics.clone(),
89 contract: None,
90 body,
91 define_opaque: None,
92 eii_impls: Default::default(),
93 });
94
95 let extern_gpu_kernel = ast::Extern::from_abi(
96 Some(ast::StrLit {
97 symbol: sym::gpu_kernel,
98 suffix: None,
99 symbol_unescaped: sym::gpu_kernel,
100 style: ast::StrStyle::Cooked,
101 span,
102 }),
103 span,
104 );
105 device_fn.sig.header.ext = extern_gpu_kernel;
106 device_fn.sig.header.safety = ast::Safety::Unsafe(span);
107
108 let rustc_offload_kernel_attr =
110 Box::new(ast::NormalAttr::from_ident(Ident::with_dummy_span(sym::rustc_offload_kernel)));
111 let rustc_offload_kernel = outer_normal_attr(
112 &rustc_offload_kernel_attr,
113 ecx.sess.psess.attr_id_generator.mk_attr_id(),
114 span,
115 );
116
117 let unsafe_item = AttrItem {
119 unsafety: ast::Safety::Unsafe(span),
120 path: ast::Path::from_ident(Ident::new(sym::no_mangle, span)),
121 args: ast::AttrItemKind::Unparsed(ast::AttrArgs::Empty),
122 };
123
124 let no_mangle_attr = Box::new(ast::NormalAttr { item: unsafe_item, tokens: None });
125 let new_id = ecx.sess.psess.attr_id_generator.mk_attr_id();
126 let unsafe_no_mangle = outer_normal_attr(&no_mangle_attr, new_id, span);
127
128 let device_item = {
129 let mut item = ecx.item(
130 span,
131 {
let len = [(), ()].len();
let mut vec = ::thin_vec::ThinVec::with_capacity(len);
vec.push(rustc_offload_kernel);
vec.push(unsafe_no_mangle);
vec
}thin_vec![rustc_offload_kernel, unsafe_no_mangle],
132 ast::ItemKind::Fn(device_fn),
133 );
134 item.vis = vis.clone();
135 Annotatable::Item(item)
136 };
137
138 let macro_expr = ecx.expr_macro_call(
140 span,
141 ecx.macro_call(
142 span,
143 ecx.path_global(
144 span,
145 [sym::std, sym::unimplemented].map(|s| Ident::new(s, span)).to_vec(),
146 ),
147 Delimiter::Parenthesis,
148 TokenStream::default(),
149 ),
150 );
151 let stmt = ecx.stmt_expr(macro_expr);
152 let body = ecx.block(span, {
let len = [()].len();
let mut vec = ::thin_vec::ThinVec::with_capacity(len);
vec.push(stmt);
vec
}thin_vec![stmt]);
153
154 let mut host_fn = Box::new(ast::Fn {
156 defaultness: ast::Defaultness::Implicit,
157 sig: sig.clone(),
158 ident,
159 generics: generics.clone(),
160 contract: None,
161 body: Some(body),
162 define_opaque: None,
163 eii_impls: Default::default(),
164 });
165
166 for param in host_fn.sig.decl.inputs.iter_mut() {
167 param.pat = Box::new(ecx.pat_wild(param.pat.span));
168 }
169
170 let ts: Vec<TokenTree> = ::alloc::boxed::box_assume_init_into_vec_unsafe(::alloc::intrinsics::write_box_via_move(::alloc::boxed::Box::new_uninit(),
[TokenTree::Token(Token::new(TokenKind::Ident(sym::never,
false.into()), span), Spacing::Joint)]))vec![TokenTree::Token(
172 Token::new(TokenKind::Ident(sym::never, false.into()), span),
173 Spacing::Joint,
174 )];
175
176 let never_arg = ast::DelimArgs {
177 dspan: DelimSpan::from_single(span),
178 delim: Delimiter::Parenthesis,
179 tokens: TokenStream::from_iter(ts),
180 };
181
182 let inline_item = ast::AttrItem {
183 unsafety: ast::Safety::Default,
184 path: ast::Path::from_ident(Ident::with_dummy_span(sym::inline)),
185 args: rustc_ast::ast::AttrItemKind::Unparsed(ast::AttrArgs::Delimited(never_arg)),
186 };
187 let inline_never_attr = Box::new(ast::NormalAttr { item: inline_item, tokens: None });
188
189 let new_id = ecx.sess.psess.attr_id_generator.mk_attr_id();
190 let inline_never = outer_normal_attr(&inline_never_attr, new_id, span);
191
192 let new_id = ecx.sess.psess.attr_id_generator.mk_attr_id();
193 let unsafe_no_mangle = outer_normal_attr(&no_mangle_attr, new_id, span);
194
195 let host_item = {
196 let mut item =
197 ecx.item(span, {
let len = [(), ()].len();
let mut vec = ::thin_vec::ThinVec::with_capacity(len);
vec.push(unsafe_no_mangle);
vec.push(inline_never);
vec
}thin_vec![unsafe_no_mangle, inline_never], ast::ItemKind::Fn(host_fn));
198 item.vis = vis.clone();
199 Annotatable::Item(item)
200 };
201
202 if compile_for_device(ecx) { ::alloc::boxed::box_assume_init_into_vec_unsafe(::alloc::intrinsics::write_box_via_move(::alloc::boxed::Box::new_uninit(),
[device_item]))vec![device_item] } else { ::alloc::boxed::box_assume_init_into_vec_unsafe(::alloc::intrinsics::write_box_via_move(::alloc::boxed::Box::new_uninit(),
[host_item]))vec![host_item] }
203}