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 = ¶m.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}