Skip to main content

cartel_gen/
lib.rs

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