Skip to main content

cartel_gen/
lib.rs

1#![allow(clippy::too_many_arguments)]
2
3extern crate proc_macro;
4
5use backend::Backend;
6use proc_macro::TokenStream;
7use syn::Token;
8use syn::punctuated::Punctuated;
9use syn::{Data, DeriveInput, Fields, ItemImpl, parse_macro_input, parse_quote};
10
11mod backend;
12mod build;
13mod derive_table;
14mod emit;
15mod parse;
16mod pg;
17mod row_meta;
18mod shape;
19mod sqlite;
20mod util;
21mod where_clause;
22
23#[proc_macro_attribute]
24pub fn dispatcher(attr: TokenStream, item: TokenStream) -> TokenStream {
25    let public_constructor = if attr.is_empty() {
26        false
27    } else {
28        let option = parse_macro_input!(attr as syn::Ident);
29        if option != "new" {
30            return syn::Error::new_spanned(option, "expected `new`")
31                .to_compile_error()
32                .into();
33        }
34        true
35    };
36
37    let mut input = parse_macro_input!(item as DeriveInput);
38    if let Err(error) = reject_packed(&input.attrs) {
39        return error.to_compile_error().into();
40    }
41    let name = input.ident.clone();
42    let vis = input.vis.clone();
43    let generics = input.generics.clone();
44    let fields = match &mut input.data {
45        Data::Struct(data) => match &mut data.fields {
46            Fields::Named(fields) => &mut fields.named,
47            _ => {
48                return syn::Error::new_spanned(
49                    &input.ident,
50                    "#[dispatcher] requires named fields",
51                )
52                .to_compile_error()
53                .into();
54            }
55        },
56        _ => {
57            return syn::Error::new_spanned(&input.ident, "#[dispatcher] requires a struct")
58                .to_compile_error()
59                .into();
60        }
61    };
62    let constructor_fields = fields
63        .iter()
64        .map(|field| {
65            let name = field
66                .ident
67                .clone()
68                .ok_or_else(|| syn::Error::new_spanned(field, "expected a named field"))?;
69            Ok((
70                name,
71                field.ty.clone(),
72                field
73                    .attrs
74                    .iter()
75                    .filter(|attr| attr.path().is_ident("cfg"))
76                    .cloned()
77                    .collect::<Vec<_>>(),
78            ))
79        })
80        .collect::<syn::Result<Vec<_>>>();
81    let constructor_fields = match constructor_fields {
82        Ok(fields) => fields,
83        Err(error) => return error.to_compile_error().into(),
84    };
85    let has_manifold_fields = fields.iter().any(|field| {
86        field
87            .attrs
88            .iter()
89            .any(|attr| attr.path().is_ident("manifold"))
90    });
91    for field in fields.iter_mut() {
92        if !has_manifold_fields {
93            field.attrs.push(parse_quote!(#[manifold]));
94        }
95        if field
96            .attrs
97            .iter()
98            .any(|attr| attr.path().is_ident("manifold"))
99        {
100            field.attrs.push(parse_quote!(#[pin]));
101        }
102    }
103    let brand = generics.lifetimes().next().map(|param| {
104        let lifetime = &param.lifetime;
105        let mut field_name = "__cartel_dispatcher_brand".to_owned();
106        while fields.iter().any(|field| {
107            field
108                .ident
109                .as_ref()
110                .is_some_and(|ident| ident == field_name.as_str())
111        }) {
112            field_name.push('_');
113        }
114        let field_name = syn::Ident::new(&field_name, proc_macro2::Span::call_site());
115        fields.push(parse_quote! {
116            #field_name: ::core::marker::PhantomData<&#lifetime ()>
117        });
118        quote::quote! { #field_name: ::core::marker::PhantomData, }
119    });
120    let constructor_args = constructor_fields
121        .iter()
122        .map(|(name, ty, attrs)| quote::quote!(#(#attrs)* #name: #ty))
123        .collect::<Vec<_>>();
124    let field_initializers = constructor_fields
125        .iter()
126        .map(|(name, _, attrs)| quote::quote!(#(#attrs)* #name,))
127        .collect::<Vec<_>>();
128    let constructor_names = constructor_fields
129        .iter()
130        .map(|(name, _, attrs)| quote::quote!(#(#attrs)* #name))
131        .collect::<Vec<_>>();
132    let (impl_generics, ty_generics, where_clause) = generics.split_for_impl();
133    let new = public_constructor.then(|| {
134        quote::quote! {
135            #[inline(always)]
136            #vis fn new(#(#constructor_args),*) -> Self {
137                Self::__cartel_dispatcher_new(#(#constructor_names),*)
138            }
139        }
140    });
141
142    quote::quote! {
143        #[::cartel_core::__private::pin_project]
144        #[derive(::cartel_core::__private::Dispatcher)]
145        #input
146
147        impl #impl_generics #name #ty_generics #where_clause {
148            #[doc(hidden)]
149            #[inline(always)]
150            fn __cartel_dispatcher_new(#(#constructor_args),*) -> Self {
151                Self {
152                    #(#field_initializers)*
153                    #brand
154                }
155            }
156
157            #new
158        }
159    }
160    .into()
161}
162
163fn reject_packed(attrs: &[syn::Attribute]) -> syn::Result<()> {
164    for attr in attrs {
165        if !attr.path().is_ident("repr") {
166            continue;
167        }
168        let reprs = attr.parse_args_with(Punctuated::<syn::Meta, Token![,]>::parse_terminated)?;
169        if let Some(repr) = reprs.iter().find(|repr| match repr {
170            syn::Meta::Path(path) => path.is_ident("packed"),
171            syn::Meta::List(list) => list.path.is_ident("packed"),
172            syn::Meta::NameValue(value) => value.path.is_ident("packed"),
173        }) {
174            return Err(syn::Error::new_spanned(
175                repr,
176                "pinned projection does not support repr(packed)",
177            ));
178        }
179    }
180    Ok(())
181}
182
183#[proc_macro_derive(PgTable, attributes(pk, table_name))]
184pub fn pg_derive_table(input: TokenStream) -> TokenStream {
185    let input = parse_macro_input!(input as DeriveInput);
186    match pg::PgBackend::derive_table(&input) {
187        Ok(ts) => ts.into(),
188        Err(e) => e.to_compile_error().into(),
189    }
190}
191
192#[proc_macro_attribute]
193pub fn query_group(_attr: TokenStream, item: TokenStream) -> TokenStream {
194    let block = parse_macro_input!(item as ItemImpl);
195    match pg::PgBackend::expand_query_group(block) {
196        Ok(ts) => ts.into(),
197        Err(e) => e.to_compile_error().into(),
198    }
199}
200
201#[proc_macro]
202pub fn pg_instance(input: TokenStream) -> TokenStream {
203    match pg::PgBackend::expand_instance(input) {
204        Ok(ts) => ts.into(),
205        Err(e) => e.to_compile_error().into(),
206    }
207}
208
209#[proc_macro_derive(SqliteTable, attributes(pk, table_name))]
210pub fn sqlite_derive_table(input: TokenStream) -> TokenStream {
211    let input = parse_macro_input!(input as DeriveInput);
212    match sqlite::SqliteBackend::derive_table(&input) {
213        Ok(ts) => ts.into(),
214        Err(e) => e.to_compile_error().into(),
215    }
216}
217
218#[proc_macro_attribute]
219pub fn sqlite_query(attr: TokenStream, item: TokenStream) -> TokenStream {
220    let f = parse_macro_input!(item as syn::ItemFn);
221    let no_probe = attr.to_string().contains("no_probe");
222    match sqlite::SqliteBackend::query_free(f, no_probe) {
223        Ok(ts) => ts.into(),
224        Err(e) => e.to_compile_error().into(),
225    }
226}