1use 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#[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 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!("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 #[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 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 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 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 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 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 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 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 #[test]
474 fn warm_test_args_default_values_when_omitted() {
475 let args: WarmTestArgs = syn::parse_str("").unwrap();
476 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}