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