1#![allow(clippy::large_enum_variant)]
2
3use core::panic;
4
5use error::error_into_token_stream;
6use generate::autotune::generate_autotune_key;
7use parse::{
8 cube_impl::CubeImpl,
9 cube_trait::{CubeTrait, CubeTraitImpl},
10 helpers::RemoveHelpers,
11 kernel::{Launch, from_tokens},
12};
13use proc_macro::TokenStream;
14use quote::quote;
15use syn::{Item, visit_mut::VisitMut};
16
17use crate::{
18 generate::{assign::generate_cube_type_mut, into_runtime::generate_into_runtime},
19 parse::{
20 cube_type::generate_cube_type, derive_expand::generate_derive_expand,
21 helpers::ReplaceDefines,
22 },
23};
24
25mod error;
26mod expression;
27mod generate;
28mod operator;
29mod parse;
30mod paths;
31mod scope;
32mod statement;
33
34#[proc_macro_attribute]
56pub fn cube(args: TokenStream, input: TokenStream) -> TokenStream {
57 match cube_impl(args, input.clone()) {
58 Ok(tokens) => tokens,
59 Err(e) => error_into_token_stream(e, input.into()).into(),
60 }
61}
62
63fn cube_impl(args: TokenStream, input: TokenStream) -> syn::Result<TokenStream> {
64 let mut item: Item = syn::parse(input)?;
65 let args = from_tokens(args.into())?;
66
67 let tokens = match item.clone() {
68 Item::Fn(kernel) => {
69 let kernel = Launch::from_item_fn(kernel, args)?;
70 RemoveHelpers.visit_item_mut(&mut item);
71 ReplaceDefines.visit_item_mut(&mut item);
72
73 let extra_allow = match kernel.func.context.is_intrinsic {
74 true => quote![#[allow(unused_variables)]],
75 false => quote![],
76 };
77
78 return Ok(TokenStream::from(quote! {
79 #[allow(dead_code, clippy::too_many_arguments)]
80 #extra_allow
81 #item
82 #kernel
83 }));
84 }
85 Item::Trait(kernel_trait) => {
86 let is_debug = args.debug.is_present();
87 let expand_trait = CubeTrait::from_item_trait(kernel_trait, args)?;
88
89 let tokens = TokenStream::from(quote! {
90 #expand_trait
91 });
92 if is_debug {
93 panic!("{tokens}");
94 }
95 return Ok(tokens);
96 }
97 Item::Impl(item_impl) => {
98 if item_impl.trait_.is_some() {
99 let mut expand_impl = CubeTraitImpl::from_item_impl(item_impl, &args)?;
100 let expand_impl = expand_impl.to_tokens_mut();
101
102 Ok(TokenStream::from(quote! {
103 #expand_impl
104 }))
105 } else {
106 let mut expand_impl = CubeImpl::from_item_impl(item_impl, &args)?;
107 let expand_impl = expand_impl.to_tokens_mut();
108
109 Ok(TokenStream::from(quote! {
110 #expand_impl
111 }))
112 }
113 }
114 item => Err(syn::Error::new_spanned(
115 item,
116 "`#[cube]` is only supported on traits and functions",
117 ))?,
118 };
119
120 if args.debug.is_present() {
121 match tokens {
122 Ok(tokens) => panic!("{tokens}"),
123 Err(err) => panic!("{err}"),
124 };
125 }
126
127 tokens
128}
129
130#[proc_macro_derive(CubeLaunch, attributes(cube, launch))]
132pub fn module_derive_cube_launch(input: TokenStream) -> TokenStream {
133 gen_cube_type(input, true)
134}
135
136#[proc_macro_derive(CubeType, attributes(cube, expand))]
138pub fn module_derive_cube_type(input: TokenStream) -> TokenStream {
139 gen_cube_type(input, false)
140}
141
142fn gen_cube_type(input: TokenStream, with_launch: bool) -> TokenStream {
143 let parsed = syn::parse(input);
144
145 let input = match &parsed {
146 Ok(val) => val,
147 Err(err) => return err.to_compile_error().into(),
148 };
149
150 match generate_cube_type(input, with_launch) {
151 Ok(val) => val.into(),
152 Err(err) => err.to_compile_error().into(),
153 }
154}
155
156#[proc_macro_attribute]
159pub fn derive_cube_comptime(_metadata: TokenStream, input: TokenStream) -> TokenStream {
160 let input: proc_macro2::TokenStream = input.into();
161 quote! {
162 #[derive(Debug, Hash, PartialEq, Eq, Clone, Copy)]
163 #input
164 }
165 .into()
166}
167
168#[proc_macro_attribute]
170pub fn derive_expand(metadata: TokenStream, input: TokenStream) -> TokenStream {
171 match generate_derive_expand(input.into(), metadata.into()) {
172 Ok(val) => val.into(),
173 Err(err) => err.to_compile_error().into(),
174 }
175}
176
177#[proc_macro]
191pub fn comptime(input: TokenStream) -> TokenStream {
192 let tokens: proc_macro2::TokenStream = input.into();
193 quote![{ #tokens }].into()
194}
195
196#[proc_macro]
209pub fn intrinsic(_input: TokenStream) -> TokenStream {
210 quote![{ cubecl::unexpanded!() }].into()
211}
212
213#[proc_macro]
228pub fn comptime_type(input: TokenStream) -> TokenStream {
229 let tokens: proc_macro2::TokenStream = input.into();
230 quote![ #tokens ].into()
231}
232
233#[proc_macro]
245pub fn comment(input: TokenStream) -> TokenStream {
246 let tokens: proc_macro2::TokenStream = input.into();
247 quote![{ #tokens }].into()
248}
249
250#[proc_macro]
266pub fn terminate(input: TokenStream) -> TokenStream {
267 let tokens: proc_macro2::TokenStream = input.into();
268 quote![{ #tokens }].into()
269}
270
271#[proc_macro_derive(AutotuneKey, attributes(autotune))]
296pub fn derive_autotune_key(input: TokenStream) -> TokenStream {
297 let input = syn::parse(input).unwrap();
298 match generate_autotune_key(input) {
299 Ok(tokens) => tokens.into(),
300 Err(e) => e.into_compile_error().into(),
301 }
302}
303
304#[proc_macro_derive(IntoRuntime, attributes(cube))]
306pub fn derive_into_runtime(input: TokenStream) -> TokenStream {
307 let input = syn::parse(input).unwrap();
308 match generate_into_runtime(&input) {
309 Ok(tokens) => tokens.into(),
310 Err(e) => e.into_compile_error().into(),
311 }
312}
313
314#[proc_macro_derive(CubeTypeMut, attributes(cube))]
316pub fn derive_assign(input: TokenStream) -> TokenStream {
317 let input = syn::parse(input).unwrap();
318 match generate_cube_type_mut(&input) {
319 Ok(tokens) => tokens.into(),
320 Err(e) => e.into_compile_error().into(),
321 }
322}