Skip to main content

rullst_orm_macros/
lib.rs

1extern crate proc_macro;
2
3use proc_macro::TokenStream;
4use syn::{DeriveInput, parse_macro_input};
5
6mod builder;
7mod enums;
8mod factory_observer;
9mod models;
10mod parser;
11mod privacy;
12mod relationships;
13
14#[cfg_attr(test, mutants::skip)]
15#[proc_macro_attribute]
16pub fn test(_args: TokenStream, input: TokenStream) -> TokenStream {
17    let input_fn = parse_macro_input!(input as syn::ItemFn);
18    let fn_name = &input_fn.sig.ident;
19    let vis = &input_fn.vis;
20    let block = &input_fn.block;
21    let attrs = &input_fn.attrs;
22
23    let expanded = quote::quote! {
24        #(#attrs)*
25        #[::tokio::test]
26        #vis async fn #fn_name() {
27            // Ensure DB is initialized (if already initialized in parallel, it ignores the error)
28            let _ = ::rullst_orm::Orm::init(&::std::env::var("DATABASE_URL").unwrap_or_else(|_| "sqlite::memory:".to_string())).await;
29
30            // Start transaction for Sandbox isolation
31            let tx = ::rullst_orm::Orm::begin_transaction().await.expect("Failed to begin sandbox transaction");
32            let tx_arc = ::std::sync::Arc::new(::tokio::sync::Mutex::new(Some(tx)));
33
34            // Scope the transaction globally for this tokio task
35            ::rullst_orm::CURRENT_TX.scope(tx_arc.clone(), async move {
36                // Execute user's test
37                let __test_closure = async move {
38                    #block
39                };
40                __test_closure.await;
41            }).await;
42
43            // Automatic Rollback
44            if let Some(tx) = tx_arc.lock().await.take() {
45                let _ = tx.rollback().await;
46            }
47        }
48    };
49    TokenStream::from(expanded)
50}
51
52#[cfg_attr(test, mutants::skip)]
53#[proc_macro_derive(PersonalData, attributes(privacy))]
54pub fn derive_personal_data(input: TokenStream) -> TokenStream {
55    privacy::derive_personal_data_impl(input)
56}
57
58#[cfg_attr(test, mutants::skip)]
59#[proc_macro_derive(Enum)]
60pub fn derive_enum(input: TokenStream) -> TokenStream {
61    enums::derive_enum_impl(input.into()).into()
62}
63
64#[cfg_attr(test, mutants::skip)]
65#[proc_macro_derive(Orm, attributes(orm, sqlx))]
66pub fn rullst_macro(input: TokenStream) -> TokenStream {
67    let input = parse_macro_input!(input as DeriveInput);
68
69    // Parse the input
70    let parsed = match parser::parse(&input) {
71        Ok(p) => p,
72        Err(e) => return TokenStream::from(e.to_compile_error()),
73    };
74
75    // Generate relationships
76    let rels = relationships::generate(&parsed);
77
78    // Generate the builder
79    let builder_code = builder::generate(
80        &parsed,
81        &rels.flags,
82        &rels.inits,
83        &rels.methods,
84        &rels.eager_loads,
85    );
86
87    // Generate factory and observers
88    let factory_observer_code = factory_observer::generate(&parsed);
89
90    // Generate the model impl
91    let model_code = models::generate(&parsed, &rels.model_methods);
92
93    // Combine
94    let expanded = quote::quote! {
95        #builder_code
96        #factory_observer_code
97        #model_code
98    };
99
100    TokenStream::from(expanded)
101}
102
103#[cfg(test)]
104mod tests {
105    use super::*;
106    use ::core::prelude::v1::test;
107    use syn::parse_quote;
108
109    fn run_macro_generator(input: &DeriveInput) -> (parser::ParsedModel, String, String) {
110        let parsed = parser::parse(input).unwrap();
111        let rels = relationships::generate(&parsed);
112        let builder = builder::generate(
113            &parsed,
114            &rels.flags,
115            &rels.inits,
116            &rels.methods,
117            &rels.eager_loads,
118        );
119        let _factory = factory_observer::generate(&parsed);
120        let models = models::generate(&parsed, &rels.model_methods);
121        (parsed, builder.to_string(), models.to_string())
122    }
123
124    #[test]
125    fn test_basic_model() {
126        let input: DeriveInput = parse_quote! {
127            #[derive(Orm)]
128            #[orm(table = "users", searchable)]
129            pub struct User {
130                pub id: i32,
131                pub name: String,
132                pub email: String,
133            }
134        };
135        let (parsed, builder, models) = run_macro_generator(&input);
136        assert_eq!(parsed.table_name, "users");
137        assert!(builder.contains("where_id"));
138        assert!(models.contains("fn delete"));
139        assert!(models.contains("fn search"));
140    }
141
142    #[test]
143    fn test_model_with_relations() {
144        let input: DeriveInput = parse_quote! {
145            #[derive(Orm)]
146            pub struct Post {
147                pub id: i32,
148                pub title: String,
149                #[orm(has_many = "Comment", foreign_key = "post_id", local_key = "id")]
150                comments: Option<Vec<Comment>>,
151                #[orm(has_one = "Author", foreign_key = "post_id", local_key = "id")]
152                author: Option<Author>,
153                #[orm(belongs_to = "User", foreign_key = "user_id", local_key = "id")]
154                user: Option<User>,
155                #[orm(belongs_to_many = "Tag", pivot_table = "post_tags", foreign_key = "post_id", related_key = "tag_id")]
156                tags: Option<Vec<Tag>>,
157                #[orm(morph_one = "Image", morph_name = "imageable")]
158                image: Option<Image>,
159                #[orm(morph_many = "Comment", morph_name = "commentable")]
160                morph_comments: Option<Vec<Comment>>,
161            }
162        };
163        let (parsed, _, _) = run_macro_generator(&input);
164        assert!(!parsed.relations.is_empty());
165    }
166
167    #[test]
168    fn test_model_with_soft_deletes() {
169        let input: DeriveInput = parse_quote! {
170            #[derive(Orm)]
171            pub struct User {
172                pub id: i32,
173                pub name: String,
174                pub deleted_at: Option<String>,
175            }
176        };
177        let (parsed, builder, _) = run_macro_generator(&input);
178        assert!(parsed.has_soft_deletes);
179        assert!(builder.contains("deleted_at IS NULL"));
180    }
181
182    #[test]
183    fn test_model_with_hidden_fields() {
184        let input: DeriveInput = parse_quote! {
185            #[derive(Orm)]
186            pub struct User {
187                pub id: i32,
188                pub name: String,
189                #[orm(hidden)]
190                pub password: String,
191            }
192        };
193        let (parsed, _, models) = run_macro_generator(&input);
194        assert_eq!(parsed.hidden_fields.len(), 1);
195        assert!(models.contains("password"));
196    }
197
198    #[test]
199    fn test_model_with_explicit_soft_delete_config() {
200        let input: DeriveInput = parse_quote! {
201            #[derive(Orm)]
202            #[orm(soft_delete(field = "is_deleted", value = "0", delval = "1"))]
203            pub struct Post {
204                pub id: i32,
205                pub title: String,
206                pub is_deleted: i32,
207            }
208        };
209        let (parsed, _, _) = run_macro_generator(&input);
210        assert!(parsed.has_soft_deletes);
211    }
212
213    #[test]
214    fn test_model_with_all_hooks_and_scopes() {
215        let input: DeriveInput = parse_quote! {
216            #[derive(Orm)]
217            #[orm(global_scope = "active", tenant_column = "account_id", before_save = "hash_pwd", after_save = "log_evt", before_delete = "check_perm", after_delete = "clear_cache", after_fetch = "decrypt_data")]
218            pub struct User {
219                pub id: i32,
220            }
221        };
222        let (parsed, _, _) = run_macro_generator(&input);
223        assert_eq!(parsed.global_scope, "active");
224        assert_eq!(parsed.tenant_column, "account_id");
225        assert_eq!(parsed.before_save, "hash_pwd");
226        assert_eq!(parsed.after_save, "log_evt");
227        assert_eq!(parsed.before_delete, "check_perm");
228        assert_eq!(parsed.after_delete, "clear_cache");
229        assert_eq!(parsed.after_fetch, "decrypt_data");
230    }
231
232    #[test]
233    fn test_model_with_soft_delete_null_sentinel() {
234        let input: DeriveInput = parse_quote! {
235            #[derive(Orm)]
236            #[orm(soft_delete(field = "deleted_at", value = "null", delval = "now()"))]
237            pub struct Audit {
238                pub id: i32,
239                pub message: String,
240                pub deleted_at: Option<String>,
241            }
242        };
243        run_macro_generator(&input);
244    }
245
246    #[test]
247    fn test_model_with_soft_delete_bigint_timestamp() {
248        let input: DeriveInput = parse_quote! {
249            #[derive(Orm)]
250            #[orm(soft_delete(field = "deleted_at", value = "0", delval = "UNIX_TIMESTAMP()"))]
251            pub struct Article {
252                pub id: i32,
253                pub title: String,
254                pub deleted_at: i64,
255            }
256        };
257        run_macro_generator(&input);
258    }
259
260    #[test]
261    fn test_model_with_orm_skip_field() {
262        let input: DeriveInput = parse_quote! {
263            #[derive(Orm)]
264            pub struct Account {
265                pub id: i32,
266                pub name: String,
267                #[orm(skip)]
268                pub password_hash: String,
269            }
270        };
271        run_macro_generator(&input);
272    }
273
274    #[test]
275    fn test_model_with_sqlx_skip_field() {
276        let input: DeriveInput = parse_quote! {
277            #[derive(Orm)]
278            pub struct Account {
279                pub id: i32,
280                pub name: String,
281                #[sqlx(skip)]
282                pub password_hash: String,
283            }
284        };
285        run_macro_generator(&input);
286    }
287
288    #[test]
289    fn test_model_with_combined_soft_delete_and_skip() {
290        let input: DeriveInput = parse_quote! {
291            #[derive(Orm)]
292            #[orm(soft_delete(field = "is_active", value = "true", delval = "false"))]
293            pub struct User {
294                pub id: i32,
295                pub name: String,
296                pub is_active: bool,
297                #[sqlx(skip)]
298                pub internal_note: String,
299            }
300        };
301        run_macro_generator(&input);
302    }
303
304    #[test]
305    fn test_parser_errors() {
306        // lowercase relation model
307        let input: DeriveInput = parse_quote! {
308            #[derive(Orm)]
309            pub struct Post {
310                pub id: i32,
311                #[orm(has_many = "comment")]
312                comments: Option<Vec<Comment>>,
313            }
314        };
315        let res = parser::parse(&input);
316        println!(
317            "PARSE RESULT FOR LOWERCASE RELATION: {:?}",
318            res.as_ref().map(|p| &p.table_name)
319        );
320        assert!(res.is_err());
321
322        // empty table name
323        let input: DeriveInput = parse_quote! {
324            #[derive(Orm)]
325            #[orm(table = "")]
326            pub struct Post {
327                pub id: i32,
328            }
329        };
330        assert!(parser::parse(&input).is_err());
331
332        // empty has_many
333        let input: DeriveInput = parse_quote! {
334            #[derive(Orm)]
335            pub struct Post {
336                pub id: i32,
337                #[orm(has_many = "")]
338                comments: Option<Vec<Comment>>,
339            }
340        };
341        assert!(parser::parse(&input).is_err());
342    }
343}