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            (
66                field.ident.clone().expect("named field"),
67                field.ty.clone(),
68                field
69                    .attrs
70                    .iter()
71                    .filter(|attr| attr.path().is_ident("cfg"))
72                    .cloned()
73                    .collect::<Vec<_>>(),
74            )
75        })
76        .collect::<Vec<_>>();
77    let has_manifold_fields = fields.iter().any(|field| {
78        field
79            .attrs
80            .iter()
81            .any(|attr| attr.path().is_ident("manifold"))
82    });
83    for field in fields.iter_mut() {
84        if !has_manifold_fields {
85            field.attrs.push(parse_quote!(#[manifold]));
86        }
87        if field
88            .attrs
89            .iter()
90            .any(|attr| attr.path().is_ident("manifold"))
91        {
92            field.attrs.push(parse_quote!(#[pin]));
93        }
94    }
95    let brand = generics.lifetimes().next().map(|param| {
96        let lifetime = &param.lifetime;
97        let mut field_name = "__cartel_dispatcher_brand".to_owned();
98        while fields.iter().any(|field| {
99            field
100                .ident
101                .as_ref()
102                .is_some_and(|ident| ident == field_name.as_str())
103        }) {
104            field_name.push('_');
105        }
106        let field_name = syn::Ident::new(&field_name, proc_macro2::Span::call_site());
107        fields.push(parse_quote! {
108            #field_name: ::core::marker::PhantomData<&#lifetime ()>
109        });
110        quote::quote! { #field_name: ::core::marker::PhantomData, }
111    });
112    let constructor_args = constructor_fields
113        .iter()
114        .map(|(name, ty, attrs)| quote::quote!(#(#attrs)* #name: #ty))
115        .collect::<Vec<_>>();
116    let field_initializers = constructor_fields
117        .iter()
118        .map(|(name, _, attrs)| quote::quote!(#(#attrs)* #name,))
119        .collect::<Vec<_>>();
120    let constructor_names = constructor_fields
121        .iter()
122        .map(|(name, _, attrs)| quote::quote!(#(#attrs)* #name))
123        .collect::<Vec<_>>();
124    let (impl_generics, ty_generics, where_clause) = generics.split_for_impl();
125    let new = public_constructor.then(|| {
126        quote::quote! {
127            #[inline(always)]
128            #vis fn new(#(#constructor_args),*) -> Self {
129                Self::__cartel_dispatcher_new(#(#constructor_names),*)
130            }
131        }
132    });
133
134    quote::quote! {
135        #[::cartel_core::__private::pin_project]
136        #[derive(::cartel_core::__private::Dispatcher)]
137        #input
138
139        impl #impl_generics #name #ty_generics #where_clause {
140            #[doc(hidden)]
141            #[inline(always)]
142            fn __cartel_dispatcher_new(#(#constructor_args),*) -> Self {
143                Self {
144                    #(#field_initializers)*
145                    #brand
146                }
147            }
148
149            #new
150        }
151    }
152    .into()
153}
154
155fn reject_packed(attrs: &[syn::Attribute]) -> syn::Result<()> {
156    for attr in attrs {
157        if !attr.path().is_ident("repr") {
158            continue;
159        }
160        let reprs = attr.parse_args_with(Punctuated::<syn::Meta, Token![,]>::parse_terminated)?;
161        if let Some(repr) = reprs.iter().find(|repr| match repr {
162            syn::Meta::Path(path) => path.is_ident("packed"),
163            syn::Meta::List(list) => list.path.is_ident("packed"),
164            syn::Meta::NameValue(value) => value.path.is_ident("packed"),
165        }) {
166            return Err(syn::Error::new_spanned(
167                repr,
168                "pinned projection does not support repr(packed)",
169            ));
170        }
171    }
172    Ok(())
173}
174
175#[proc_macro_derive(PgTable, attributes(pk, table_name))]
176pub fn pg_derive_table(input: TokenStream) -> TokenStream {
177    let input = parse_macro_input!(input as DeriveInput);
178    match pg::PgBackend::derive_table(&input) {
179        Ok(ts) => ts.into(),
180        Err(e) => e.to_compile_error().into(),
181    }
182}
183
184#[proc_macro_attribute]
185pub fn query_group(_attr: TokenStream, item: TokenStream) -> TokenStream {
186    let block = parse_macro_input!(item as ItemImpl);
187    match pg::PgBackend::expand_query_group(block) {
188        Ok(ts) => ts.into(),
189        Err(e) => e.to_compile_error().into(),
190    }
191}
192
193#[proc_macro]
194pub fn pg_instance(input: TokenStream) -> TokenStream {
195    match pg::PgBackend::expand_instance(input) {
196        Ok(ts) => ts.into(),
197        Err(e) => e.to_compile_error().into(),
198    }
199}
200
201#[proc_macro_derive(SqliteTable, attributes(pk, table_name))]
202pub fn sqlite_derive_table(input: TokenStream) -> TokenStream {
203    let input = parse_macro_input!(input as DeriveInput);
204    match sqlite::SqliteBackend::derive_table(&input) {
205        Ok(ts) => ts.into(),
206        Err(e) => e.to_compile_error().into(),
207    }
208}
209
210#[proc_macro_attribute]
211pub fn sqlite_query(attr: TokenStream, item: TokenStream) -> TokenStream {
212    let f = parse_macro_input!(item as syn::ItemFn);
213    let no_probe = attr.to_string().contains("no_probe");
214    match sqlite::SqliteBackend::query_free(f, no_probe) {
215        Ok(ts) => ts.into(),
216        Err(e) => e.to_compile_error().into(),
217    }
218}