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