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::{
19 asm::generate_asm_unexpanded, assign::generate_cube_type_mut,
20 into_runtime::generate_into_runtime,
21 },
22 parse::{
23 cube_type::generate_cube_type, derive_expand::generate_derive_expand,
24 helpers::ReplaceDefines,
25 },
26};
27
28mod error;
29mod expression;
30mod generate;
31mod operator;
32mod parse;
33mod paths;
34mod scope;
35mod statement;
36
37#[proc_macro_attribute]
59pub fn cube(args: TokenStream, input: TokenStream) -> TokenStream {
60 match cube_impl(args, input.clone()) {
61 Ok(tokens) => tokens,
62 Err(e) => error_into_token_stream(e, input.into()).into(),
63 }
64}
65
66fn cube_impl(args: TokenStream, input: TokenStream) -> syn::Result<TokenStream> {
67 let mut item: Item = syn::parse(input)?;
68 let args = from_tokens(args.into())?;
69
70 let tokens = match item.clone() {
71 Item::Fn(kernel) => {
72 let kernel = Launch::from_item_fn(kernel, args)?;
73 RemoveHelpers.visit_item_mut(&mut item);
74 ReplaceDefines.visit_item_mut(&mut item);
75
76 let extra_allow = match kernel.func.context.is_intrinsic {
77 true => quote![#[allow(unused_variables)]],
78 false => quote![],
79 };
80
81 return Ok(TokenStream::from(quote! {
82 #[allow(dead_code, clippy::too_many_arguments)]
83 #extra_allow
84 #item
85 #kernel
86 }));
87 }
88 Item::Trait(kernel_trait) => {
89 let is_debug = args.debug.is_present();
90 let expand_trait = CubeTrait::from_item_trait(kernel_trait, args)?;
91
92 let tokens = TokenStream::from(quote! {
93 #expand_trait
94 });
95 if is_debug {
96 panic!("{tokens}");
97 }
98 return Ok(tokens);
99 }
100 Item::Impl(item_impl) => {
101 if item_impl.trait_.is_some() {
102 let mut expand_impl = CubeTraitImpl::from_item_impl(item_impl, &args)?;
103 let expand_impl = expand_impl.to_tokens_mut();
104
105 Ok(TokenStream::from(quote! {
106 #expand_impl
107 }))
108 } else {
109 let mut expand_impl = CubeImpl::from_item_impl(item_impl, &args)?;
110 let expand_impl = expand_impl.to_tokens_mut();
111
112 Ok(TokenStream::from(quote! {
113 #expand_impl
114 }))
115 }
116 }
117 item => Err(syn::Error::new_spanned(
118 item,
119 "`#[cube]` is only supported on traits and functions",
120 ))?,
121 };
122
123 if args.debug.is_present() {
124 match tokens {
125 Ok(tokens) => panic!("{tokens}"),
126 Err(err) => panic!("{err}"),
127 };
128 }
129
130 tokens
131}
132
133#[proc_macro_derive(CubeLaunch, attributes(cube, launch))]
135pub fn module_derive_cube_launch(input: TokenStream) -> TokenStream {
136 gen_cube_type(input, true)
137}
138
139#[proc_macro_derive(CubeType, attributes(cube, expand))]
141pub fn module_derive_cube_type(input: TokenStream) -> TokenStream {
142 gen_cube_type(input, false)
143}
144
145fn gen_cube_type(input: TokenStream, with_launch: bool) -> TokenStream {
146 let parsed = syn::parse(input);
147
148 let input = match &parsed {
149 Ok(val) => val,
150 Err(err) => return err.to_compile_error().into(),
151 };
152
153 match generate_cube_type(input, with_launch) {
154 Ok(val) => val.into(),
155 Err(err) => err.to_compile_error().into(),
156 }
157}
158
159#[proc_macro_attribute]
162pub fn derive_cube_comptime(_metadata: TokenStream, input: TokenStream) -> TokenStream {
163 let input: proc_macro2::TokenStream = input.into();
164 quote! {
165 #[derive(Debug, Hash, PartialEq, Eq, Clone, Copy)]
166 #input
167 }
168 .into()
169}
170
171#[proc_macro_attribute]
173pub fn derive_expand(metadata: TokenStream, input: TokenStream) -> TokenStream {
174 match generate_derive_expand(input.into(), metadata.into()) {
175 Ok(val) => val.into(),
176 Err(err) => err.to_compile_error().into(),
177 }
178}
179
180#[proc_macro]
194pub fn comptime(input: TokenStream) -> TokenStream {
195 let tokens: proc_macro2::TokenStream = input.into();
196 quote![{ #tokens }].into()
197}
198
199#[proc_macro]
212pub fn intrinsic(_input: TokenStream) -> TokenStream {
213 quote![{ cubecl::unexpanded!() }].into()
214}
215
216#[proc_macro]
234pub fn gpu_asm(input: TokenStream) -> TokenStream {
235 match generate_asm_unexpanded(input.into()) {
236 Ok(val) => val.into(),
237 Err(err) => err.to_compile_error().into(),
238 }
239}
240
241#[proc_macro]
256pub fn comptime_type(input: TokenStream) -> TokenStream {
257 let tokens: proc_macro2::TokenStream = input.into();
258 quote![ #tokens ].into()
259}
260
261#[proc_macro]
273pub fn comment(input: TokenStream) -> TokenStream {
274 let tokens: proc_macro2::TokenStream = input.into();
275 quote![{ #tokens }].into()
276}
277
278#[proc_macro]
294pub fn terminate(input: TokenStream) -> TokenStream {
295 let tokens: proc_macro2::TokenStream = input.into();
296 quote![{ #tokens }].into()
297}
298
299#[proc_macro_derive(AutotuneKey, attributes(autotune))]
324pub fn derive_autotune_key(input: TokenStream) -> TokenStream {
325 let input = syn::parse(input).unwrap();
326 match generate_autotune_key(input) {
327 Ok(tokens) => tokens.into(),
328 Err(e) => e.into_compile_error().into(),
329 }
330}
331
332#[proc_macro_derive(IntoRuntime, attributes(cube))]
334pub fn derive_into_runtime(input: TokenStream) -> TokenStream {
335 let input = syn::parse(input).unwrap();
336 match generate_into_runtime(&input) {
337 Ok(tokens) => tokens.into(),
338 Err(e) => e.into_compile_error().into(),
339 }
340}
341
342#[proc_macro_derive(CubeTypeMut, attributes(cube))]
344pub fn derive_assign(input: TokenStream) -> TokenStream {
345 let input = syn::parse(input).unwrap();
346 match generate_cube_type_mut(&input) {
347 Ok(tokens) => tokens.into(),
348 Err(e) => e.into_compile_error().into(),
349 }
350}