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::{parse_macro_input, Expr, ExprLit, FnArg, ItemFn, Lit, MetaNameValue, Token};
10
11struct WarmTestArgs {
12    migrations: Option<String>,
13    database_url_env: Option<String>,
14}
15
16impl Parse for WarmTestArgs {
17    fn parse(input: ParseStream) -> syn::Result<Self> {
18        let mut migrations = None;
19        let mut database_url_env = None;
20
21        let pairs = Punctuated::<MetaNameValue, Token![,]>::parse_terminated(input)?;
22        for pair in pairs {
23            let key = pair
24                .path
25                .get_ident()
26                .map(|i| i.to_string())
27                .unwrap_or_default();
28
29            let value = match &pair.value {
30                Expr::Lit(ExprLit {
31                    lit: Lit::Str(s), ..
32                }) => s.value(),
33                other => {
34                    return Err(syn::Error::new_spanned(
35                        other,
36                        "expected a string literal, e.g. migrations = \"./migrations\"",
37                    ))
38                }
39            };
40
41            match key.as_str() {
42                "migrations" => migrations = Some(value),
43                "database_url_env" => database_url_env = Some(value),
44                other => {
45                    return Err(syn::Error::new_spanned(
46                        &pair.path,
47                        format!(
48                            "unknown #[warm_test] argument `{other}` \
49                             (expected `migrations` or `database_url_env`)"
50                        ),
51                    ))
52                }
53            }
54        }
55
56        Ok(WarmTestArgs {
57            migrations,
58            database_url_env,
59        })
60    }
61}
62
63///  Wraps an async test function so it receives a freshly cloned, migrated
64/// `sqlx::PgPool` cloned from a template that's built once.
65///  The test database is dropped after the function returns, whether it panics or succeeds. 
66///  The template is built once per unique migration set, and reused across every test and every run,
67///  so tests run in milliseconds regardless of schema size.
68///
69/// ```ignore
70/// #[warmpool::warm_test(migrations = "./migrations")]
71/// async fn creates_a_post(pool: sqlx::PgPool) {
72///     let row: (i64,) = sqlx::query_as("SELECT count(*) FROM posts")
73///         .fetch_one(&pool)
74///         .await
75///         .unwrap();
76///     assert_eq!(row.0, 0);
77/// }
78/// ```
79///
80/// Arguments (both optional):
81/// - `migrations = "./path"`  defaults to `./migrations`.
82/// - `database_url_env = "ENV_VAR"` defaults to `DATABASE_URL`.
83#[proc_macro_attribute]
84pub fn warm_test(attr: TokenStream, item: TokenStream) -> TokenStream {
85    let args = parse_macro_input!(attr as WarmTestArgs);
86    let input = parse_macro_input!(item as ItemFn);
87
88    let attrs = &input.attrs;
89    let vis = &input.vis;
90    let sig = &input.sig;
91    let block = &input.block;
92    let fn_name = &sig.ident;
93
94    if sig.asyncness.is_none() {
95        return syn::Error::new_spanned(sig, "#[warm_test] functions must be `async fn`")
96            .to_compile_error()
97            .into();
98    }
99
100    let pool_ident = match sig.inputs.iter().collect::<Vec<_>>().as_slice() {
101        [FnArg::Typed(pat_type)] => &pat_type.pat,
102        _ => {
103            return syn::Error::new_spanned(
104                &sig.inputs,
105                "#[warm_test] functions must take exactly one argument: \
106                 `async fn my_test(pool: sqlx::PgPool)`",
107            )
108            .to_compile_error()
109            .into()
110        }
111    };
112
113    let migrations_path = args.migrations.unwrap_or_else(|| "./migrations".to_string());
114    let database_url_env = args
115        .database_url_env
116        .unwrap_or_else(|| "DATABASE_URL".to_string());
117
118    let expanded = quote! {
119        #(#attrs)*
120        #[::tokio::test]
121        #vis async fn #fn_name() {
122            let __warmpool_database_url = ::std::env::var(#database_url_env)
123                .unwrap_or_else(|_| panic!(
124                    "warmpool: environment variable `{}` is not set", #database_url_env
125                ));
126
127            let __warmpool_connect_options: ::warmpool::__private::PgConnectOptions =
128                __warmpool_database_url.parse().unwrap_or_else(|e| panic!(
129                    "warmpool: `{}` is not a valid Postgres connection string: {}",
130                    #database_url_env, e
131                ));
132
133            let __warmpool_template = ::warmpool::TemplatePool::builder(__warmpool_connect_options)
134                .migrations_from(#migrations_path)
135                .build()
136                .await
137                .expect("warmpool: failed to build template pool");
138
139            let __warmpool_test_db = __warmpool_template
140                .create_test_database()
141                .await
142                .expect("warmpool: failed to clone test database from template");
143
144            let #pool_ident = __warmpool_test_db.pool().clone();
145
146            let __warmpool_result = ::warmpool::__private::FutureExt::catch_unwind(
147                ::warmpool::__private::AssertUnwindSafe(async #block)
148            )
149            .await;
150
151            if let Err(__warmpool_drop_err) = __warmpool_test_db.drop_database().await {
152                eprintln!("warmpool: failed to drop test database: {__warmpool_drop_err}");
153            }
154
155            if let Err(__warmpool_panic) = __warmpool_result {
156                ::std::panic::resume_unwind(__warmpool_panic);
157            }
158        }
159    };
160
161    expanded.into()
162}
163
164#[cfg(test)]
165mod tests {
166    use super::*;
167
168    #[test]
169    fn parse_no_args() {
170        let args: WarmTestArgs = syn::parse_str("").unwrap();
171        assert!(args.migrations.is_none());
172        assert!(args.database_url_env.is_none());
173    }
174
175    #[test]
176    fn parse_migrations_arg() {
177        let args: WarmTestArgs = syn::parse_str(r#"migrations = "./migrations""#).unwrap();
178        assert_eq!(args.migrations.as_deref(), Some("./migrations"));
179        assert!(args.database_url_env.is_none());
180    }
181
182    #[test]
183    fn parse_database_url_env_arg() {
184        let args: WarmTestArgs = syn::parse_str(r#"database_url_env = "TEST_DATABASE_URL""#).unwrap();
185        assert!(args.migrations.is_none());
186        assert_eq!(args.database_url_env.as_deref(), Some("TEST_DATABASE_URL"));
187    }
188
189    #[test]
190    fn parse_both_args() {
191        let args: WarmTestArgs = syn::parse_str(
192            r#"migrations = "./migrations", database_url_env = "TEST_DATABASE_URL""#,
193        )
194        .unwrap();
195
196        assert_eq!(args.migrations.as_deref(), Some("./migrations"));
197        assert_eq!(args.database_url_env.as_deref(), Some("TEST_DATABASE_URL"));
198    }
199
200    #[test]
201    fn reject_unknown_argument() {
202        assert!(syn::parse_str::<WarmTestArgs>(r#"foo = \"bar\""#).is_err());
203    }
204
205    #[test]
206    fn reject_non_string_literal_argument() {
207        assert!(syn::parse_str::<WarmTestArgs>(r#"migrations = 42"#).is_err());
208    }
209}