1use 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#[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}