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