Skip to main content

warmpool_macros/
lib.rs

1//! Proc-macro companion for `warmpool`. Depend on `warmpool` with the
2//! `macros` feature enabled rather than on this crate directly, it's
3//! re-exported from there.
4
5use proc_macro::TokenStream;
6use quote::quote;
7use syn::parse::{Parse, ParseStream};
8use syn::punctuated::Punctuated;
9use syn::{Expr, ExprLit, FnArg, ItemFn, Lit, MetaNameValue, Token, parse_macro_input};
10
11#[derive(Debug)]
12struct WarmTestArgs {
13    migrations: Option<String>,
14    database_url_env: Option<String>,
15    clone_strategy: Option<String>,
16}
17
18impl Parse for WarmTestArgs {
19    fn parse(input: ParseStream) -> syn::Result<Self> {
20        let mut migrations = None;
21        let mut database_url_env = None;
22        let mut clone_strategy = None;
23
24        let pairs = Punctuated::<MetaNameValue, Token![,]>::parse_terminated(input)?;
25        for pair in pairs {
26            let key = pair
27                .path
28                .get_ident()
29                .map(|i| i.to_string())
30                .unwrap_or_default();
31
32            let value = match &pair.value {
33                Expr::Lit(ExprLit {
34                    lit: Lit::Str(s), ..
35                }) => s.value(),
36                other => {
37                    return Err(syn::Error::new_spanned(
38                        other,
39                        "expected a string literal, e.g. migrations = \"./migrations\"",
40                    ));
41                }
42            };
43
44            match key.as_str() {
45                "migrations" => migrations = Some(value),
46                "database_url_env" => database_url_env = Some(value),
47                "clone_strategy" => {
48                    if !matches!(value.as_str(), "wal_log" | "file_copy" | "auto") {
49                        return Err(syn::Error::new_spanned(
50                            &pair.value,
51                            format!(
52                                "invalid clone_strategy `{value}` \
53                                 (expected \"wal_log\", \"file_copy\", or \"auto\")"
54                            ),
55                        ));
56                    }
57                    clone_strategy = Some(value);
58                }
59                other => {
60                    return Err(syn::Error::new_spanned(
61                        &pair.path,
62                        format!(
63                            "unknown #[warm_test] argument `{other}` \
64                             (expected `migrations`, `database_url_env`, or `clone_strategy`)"
65                        ),
66                    ));
67                }
68            }
69        }
70
71        Ok(WarmTestArgs {
72            migrations,
73            database_url_env,
74            clone_strategy,
75        })
76    }
77}
78
79///  Wraps an async test function so it receives a freshly cloned, migrated
80/// `sqlx::PgPool` cloned from a template that's built once.
81///  The test database is dropped after the function returns, whether it panics or succeeds.
82///  The template is built once per unique migration set, and reused across every test and every run,
83///  so tests run in milliseconds regardless of schema size.
84///
85/// ```ignore
86/// #[warmpool::warm_test(migrations = "./migrations")]
87/// async fn creates_a_post(pool: sqlx::PgPool) {
88///     let row: (i64,) = sqlx::query_as("SELECT count(*) FROM posts")
89///         .fetch_one(&pool)
90///         .await
91///         .unwrap();
92///     assert_eq!(row.0, 0);
93/// }
94/// ```
95///
96/// Arguments all optional:
97/// - `migrations = "./path"`  defaults to `./migrations`.
98/// - `database_url_env = "ENV_VAR"` defaults to `DATABASE_URL`.
99/// - `clone_strategy = "wal_log" | "file_copy" | "auto"` — defaults to
100///   whatever [`warmpool::TemplatePoolBuilder`] defaults to (`wal_log` as
101///   of warmpool 0.1.1+). Rejected at compile time if it isn't one of those
102///   three strings.
103#[proc_macro_attribute]
104pub fn warm_test(attr: TokenStream, item: TokenStream) -> TokenStream {
105    let args = parse_macro_input!(attr as WarmTestArgs);
106    let input = parse_macro_input!(item as ItemFn);
107
108    let attrs = &input.attrs;
109    let vis = &input.vis;
110    let sig = &input.sig;
111    let block = &input.block;
112    let fn_name = &sig.ident;
113
114    if sig.asyncness.is_none() {
115        return syn::Error::new_spanned(sig, "#[warm_test] functions must be `async fn`")
116            .to_compile_error()
117            .into();
118    }
119
120    let pool_ident = match sig.inputs.iter().collect::<Vec<_>>().as_slice() {
121        [FnArg::Typed(pat_type)] => &pat_type.pat,
122        _ => {
123            return syn::Error::new_spanned(
124                &sig.inputs,
125                "#[warm_test] functions must take exactly one argument: \
126                 `async fn my_test(pool: sqlx::PgPool)`",
127            )
128            .to_compile_error()
129            .into();
130        }
131    };
132
133    let migrations_path = args
134        .migrations
135        .unwrap_or_else(|| "./migrations".to_string());
136    let database_url_env = args
137        .database_url_env
138        .unwrap_or_else(|| "DATABASE_URL".to_string());
139
140    // `None` here means "don't call .clone_strategy(...) at all",
141    // so the builder's own default applies, exactly as if the attribute had never
142    // been extended with this argument. Only emit the call when the user
143    // asked for something explicitly.
144    let clone_strategy_call = args.clone_strategy.map(|value| {
145        let variant = match value.as_str() {
146            "wal_log" => quote! { WalLog },
147            "file_copy" => quote! { FileCopy },
148            "auto" => quote! { Auto },
149            // Unreachable: WarmTestArgs::parse already rejected anything else.
150            _ => unreachable!("clone_strategy value validated during parsing"),
151        };
152        quote! {
153            .clone_strategy(::warmpool::CloneStrategy::#variant)
154        }
155    });
156
157    let expanded = quote! {
158        #(#attrs)*
159        #[::tokio::test]
160        #vis async fn #fn_name() {
161            let __warmpool_database_url = ::std::env::var(#database_url_env)
162                .unwrap_or_else(|_| panic!(
163                    "warmpool: environment variable `{}` is not set", #database_url_env
164                ));
165
166            let __warmpool_connect_options: ::warmpool::__private::PgConnectOptions =
167                __warmpool_database_url.parse().unwrap_or_else(|e| panic!(
168                    "warmpool: `{}` is not a valid Postgres connection string: {}",
169                    #database_url_env, e
170                ));
171
172            let __warmpool_template = ::warmpool::TemplatePool::builder(__warmpool_connect_options)
173                .migrations_from(#migrations_path)
174                #clone_strategy_call
175                .build()
176                .await
177                .expect("warmpool: failed to build template pool");
178
179            let __warmpool_test_db = __warmpool_template
180                .create_test_database()
181                .await
182                .expect("warmpool: failed to clone test database from template");
183
184            let #pool_ident = __warmpool_test_db.pool().clone();
185
186            let __warmpool_result = ::warmpool::__private::FutureExt::catch_unwind(
187                ::warmpool::__private::AssertUnwindSafe(async #block)
188            )
189            .await;
190
191            if let Err(__warmpool_drop_err) = __warmpool_test_db.drop_database().await {
192                eprintln!("warmpool: failed to drop test database: {__warmpool_drop_err}");
193            }
194
195            if let Err(__warmpool_panic) = __warmpool_result {
196                ::std::panic::resume_unwind(__warmpool_panic);
197            }
198        }
199    };
200
201    expanded.into()
202}
203
204#[cfg(test)]
205mod tests {
206    use super::*;
207
208    #[test]
209    fn parse_no_args() {
210        let args: WarmTestArgs = syn::parse_str("").unwrap();
211        assert!(args.migrations.is_none());
212        assert!(args.database_url_env.is_none());
213        assert!(args.clone_strategy.is_none());
214    }
215
216    #[test]
217    fn parse_migrations_arg() {
218        let args: WarmTestArgs = syn::parse_str(r#"migrations = "./migrations""#).unwrap();
219        assert_eq!(args.migrations.as_deref(), Some("./migrations"));
220        assert!(args.database_url_env.is_none());
221    }
222
223    #[test]
224    fn parse_database_url_env_arg() {
225        let args: WarmTestArgs =
226            syn::parse_str(r#"database_url_env = "TEST_DATABASE_URL""#).unwrap();
227        assert!(args.migrations.is_none());
228        assert_eq!(args.database_url_env.as_deref(), Some("TEST_DATABASE_URL"));
229    }
230
231    #[test]
232    fn parse_both_args() {
233        let args: WarmTestArgs = syn::parse_str(
234            r#"migrations = "./migrations", database_url_env = "TEST_DATABASE_URL""#,
235        )
236        .unwrap();
237
238        assert_eq!(args.migrations.as_deref(), Some("./migrations"));
239        assert_eq!(args.database_url_env.as_deref(), Some("TEST_DATABASE_URL"));
240    }
241
242    #[test]
243    fn reject_unknown_argument() {
244        assert!(syn::parse_str::<WarmTestArgs>(r#"foo = \"bar\""#).is_err());
245    }
246
247    #[test]
248    fn reject_non_string_literal_argument() {
249        assert!(syn::parse_str::<WarmTestArgs>(r#"migrations = 42"#).is_err());
250    }
251
252    // 0.1.1
253
254    #[test]
255    fn parse_clone_strategy_wal_log() {
256        let args: WarmTestArgs = syn::parse_str(r#"clone_strategy = "wal_log""#).unwrap();
257        assert_eq!(args.clone_strategy.as_deref(), Some("wal_log"));
258    }
259
260    #[test]
261    fn parse_clone_strategy_file_copy() {
262        let args: WarmTestArgs = syn::parse_str(r#"clone_strategy = "file_copy""#).unwrap();
263        assert_eq!(args.clone_strategy.as_deref(), Some("file_copy"));
264    }
265
266    #[test]
267    fn parse_clone_strategy_auto() {
268        let args: WarmTestArgs = syn::parse_str(r#"clone_strategy = "auto""#).unwrap();
269        assert_eq!(args.clone_strategy.as_deref(), Some("auto"));
270    }
271
272    #[test]
273    fn reject_invalid_clone_strategy_value() {
274        let result: syn::Result<WarmTestArgs> = syn::parse_str(r#"clone_strategy = "fast_please""#);
275        assert!(result.is_err());
276        assert!(
277            result
278                .unwrap_err()
279                .to_string()
280                .contains("invalid clone_strategy")
281        );
282    }
283
284    #[test]
285    fn parse_all_three_args_together() {
286        let args: WarmTestArgs = syn::parse_str(
287            r#"migrations = "./migrations", database_url_env = "TEST_DATABASE_URL", clone_strategy = "file_copy""#,
288        )
289        .unwrap();
290
291        assert_eq!(args.migrations.as_deref(), Some("./migrations"));
292        assert_eq!(args.database_url_env.as_deref(), Some("TEST_DATABASE_URL"));
293        assert_eq!(args.clone_strategy.as_deref(), Some("file_copy"));
294    }
295
296    #[test]
297    fn parse_tolerates_surrounding_whitespace() {
298        let args: WarmTestArgs = syn::parse_str(r#"  migrations = "./migrations"  "#).unwrap();
299        assert_eq!(args.migrations.as_deref(), Some("./migrations"));
300    }
301
302    #[test]
303    fn parse_tolerates_whitespace_around_equals_and_commas() {
304        let args: WarmTestArgs =
305            syn::parse_str(r#"migrations="./migrations" , clone_strategy = "auto""#).unwrap();
306        assert_eq!(args.migrations.as_deref(), Some("./migrations"));
307        assert_eq!(args.clone_strategy.as_deref(), Some("auto"));
308    }
309
310    #[test]
311    fn parse_accepts_trailing_comma() {
312        // Punctuated::parse_terminated explicitly allows a trailing
313        // separator; worth locking in since #[warm_test(migrations = "x",)]
314        // is a natural thing to type after adding/removing an argument.
315        let args: WarmTestArgs = syn::parse_str(r#"migrations = "./migrations","#).unwrap();
316        assert_eq!(args.migrations.as_deref(), Some("./migrations"));
317    }
318
319    #[test]
320    fn parse_empty_string_value_is_accepted_syntactically() {
321        // WarmTestArgs currently does no non emptiness validation for
322        // `migrations` or `database_url_env` an empty path or env var
323        // name parses fine here and would only surface as a problem later,
324        // at expansion runtime (a nonsensical env var lookup, or `Migrator`
325        // failing on `""` as a path). We know this and will fix it in a future realease,
326        // but for now the parser is happy with an empty string literal.
327        let args: WarmTestArgs = syn::parse_str(r#"migrations = """#).unwrap();
328        assert_eq!(args.migrations.as_deref(), Some(""));
329    }
330
331    #[test]
332    fn parse_clone_strategy_is_case_sensitive() {
333        // Only the exact lowercase forms are accepted "WAL_LOG",
334        // "Wal_Log", etc. are all rejected.
335        for bad in ["WAL_LOG", "Wal_Log", "FILE_COPY", "AUTO", "Auto"] {
336            let result: syn::Result<WarmTestArgs> =
337                syn::parse_str(&format!(r#"clone_strategy = "{bad}""#));
338            assert!(
339                result.is_err(),
340                "expected `{bad}` to be rejected (case-sensitive match)"
341            );
342        }
343    }
344
345    #[test]
346    fn parse_duplicate_key_lets_the_last_occurrence_win_silently() {
347        // WarmTestArgs::parse has no duplicate key detection logic, so
348        // #[warm_test(migrations = "a", migrations = "b")] silently keeps
349        // "b" with no warning or error. In a future realese this will be
350        // a compile error.
351        let args: WarmTestArgs = syn::parse_str(r#"migrations = "a", migrations = "b""#).unwrap();
352        assert_eq!(args.migrations.as_deref(), Some("b"));
353    }
354
355    #[test]
356    fn parse_duplicate_clone_strategy_key_also_lets_last_win() {
357        let args: WarmTestArgs =
358            syn::parse_str(r#"clone_strategy = "wal_log", clone_strategy = "file_copy""#).unwrap();
359        assert_eq!(args.clone_strategy.as_deref(), Some("file_copy"));
360    }
361
362    #[test]
363    fn reject_malformed_missing_equals() {
364        assert!(syn::parse_str::<WarmTestArgs>(r#"migrations "./migrations""#).is_err());
365    }
366
367    #[test]
368    fn reject_malformed_missing_value() {
369        assert!(syn::parse_str::<WarmTestArgs>(r#"migrations ="#).is_err());
370    }
371
372    #[test]
373    fn reject_malformed_leading_comma() {
374        assert!(syn::parse_str::<WarmTestArgs>(r#", migrations = "./migrations""#).is_err());
375    }
376
377    #[test]
378    fn reject_malformed_double_comma() {
379        assert!(
380            syn::parse_str::<WarmTestArgs>(r#"migrations = "a",, database_url_env = "B""#).is_err()
381        );
382    }
383
384    #[test]
385    fn reject_unknown_argument_alongside_valid_ones() {
386        // A single bad key anywhere in the list rejects the whole
387        // attribute, it isn't "parse what you can."
388        let result: syn::Result<WarmTestArgs> =
389            syn::parse_str(r#"migrations = "./migrations", bogus = "x", clone_strategy = "auto""#);
390        assert!(result.is_err());
391    }
392
393    #[test]
394    fn unknown_argument_error_message_names_the_bad_key_and_the_valid_ones() {
395        let result: syn::Result<WarmTestArgs> = syn::parse_str(r#"totally_bogus = "x""#);
396        let message = result.unwrap_err().to_string();
397        assert!(message.contains("totally_bogus"), "message was: {message}");
398        assert!(message.contains("migrations"), "message was: {message}");
399        assert!(
400            message.contains("database_url_env"),
401            "message was: {message}"
402        );
403        assert!(message.contains("clone_strategy"), "message was: {message}");
404    }
405
406    #[test]
407    fn invalid_clone_strategy_error_message_names_the_bad_value_and_the_valid_ones() {
408        let result: syn::Result<WarmTestArgs> = syn::parse_str(r#"clone_strategy = "yolo""#);
409        let message = result.unwrap_err().to_string();
410        assert!(message.contains("yolo"), "message was: {message}");
411        assert!(message.contains("wal_log"), "message was: {message}");
412        assert!(message.contains("file_copy"), "message was: {message}");
413        assert!(message.contains("auto"), "message was: {message}");
414    }
415
416    #[test]
417    fn non_string_literal_error_message_is_actionable() {
418        let result: syn::Result<WarmTestArgs> = syn::parse_str(r#"migrations = 42"#);
419        let message = result.unwrap_err().to_string();
420        assert!(
421            message.contains("string literal"),
422            "message should explain what was expected: {message}"
423        );
424    }
425
426    #[test]
427    fn reject_boolean_literal_value() {
428        assert!(syn::parse_str::<WarmTestArgs>(r#"clone_strategy = true"#).is_err());
429    }
430
431    #[test]
432    fn reject_identifier_as_value_without_quotes() {
433        // writing clone_strategy = wal_log instead of
434        // clone_strategy = "wal_log". Must fail with the "expected a
435        // string literal" message.
436        let result: syn::Result<WarmTestArgs> = syn::parse_str(r#"clone_strategy = wal_log"#);
437        assert!(result.is_err());
438        assert!(result.unwrap_err().to_string().contains("string literal"));
439    }
440
441    #[test]
442    fn parse_is_order_independent() {
443        // The three arguments can appear in any order with the same result.
444        let a: WarmTestArgs = syn::parse_str(
445            r#"migrations = "./m", database_url_env = "E", clone_strategy = "auto""#,
446        )
447        .unwrap();
448        let b: WarmTestArgs = syn::parse_str(
449            r#"clone_strategy = "auto", migrations = "./m", database_url_env = "E""#,
450        )
451        .unwrap();
452        let c: WarmTestArgs = syn::parse_str(
453            r#"database_url_env = "E", clone_strategy = "auto", migrations = "./m""#,
454        )
455        .unwrap();
456
457        for args in [a, b, c] {
458            assert_eq!(args.migrations.as_deref(), Some("./m"));
459            assert_eq!(args.database_url_env.as_deref(), Some("E"));
460            assert_eq!(args.clone_strategy.as_deref(), Some("auto"));
461        }
462    }
463
464    // codegen shape: does #[warm_test] actually reject non-async fns
465    // and wrong argument counts the way the we promises? These
466    // exercise the proc-macro attribute function itself (not just
467    // WarmTestArgs::parse), using syn to build a minimal ItemFn and
468    // checking the *shape* of the generated tokens rather than trying to
469    // execute them full end-to-end expansion is covered by the
470    // `expand_check` crate instead, which actually
471    // resolves the generated code against the real `warmpool` crate.
472
473    #[test]
474    fn warm_test_args_default_values_when_omitted() {
475        let args: WarmTestArgs = syn::parse_str("").unwrap();
476        // Mirrors the defaults applied in warm_test()'s codegen path
477        // if these ever drift apart, the doc comment's claimed defaults
478        // ("./migrations", "DATABASE_URL", builder's own clone_strategy
479        // default) would silently stop matching reality.
480        assert_eq!(
481            args.migrations
482                .unwrap_or_else(|| "./migrations".to_string()),
483            "./migrations"
484        );
485        assert_eq!(
486            args.database_url_env
487                .unwrap_or_else(|| "DATABASE_URL".to_string()),
488            "DATABASE_URL"
489        );
490        assert!(
491            args.clone_strategy.is_none(),
492            "omitted clone_strategy must not synthesize a value \
493             None is what makes warm_test() skip emitting .clone_strategy(...) \
494             entirely and fall through to the builder's own default"
495        );
496    }
497}