use proc_macro::TokenStream;
use proc_macro2::Span;
use quote::quote;
use std::collections::HashSet;
use std::sync::Mutex;
use syn::{parse::Parse, parse::ParseStream, LitInt, LitStr, Token};
/// Input for the migs! macro
/// Can be either:
/// - migs!(path = "path/to/file.sql") // scope defaults to "init", order auto-assigned
/// - migs!(path = "path/to/file.sql", scope = "...")
/// - migs!(path = "path/to/file.sql", scope = "...", order = 1)
/// - migs!(sql = "SELECT * FROM {{table}}") // scope defaults to "init", order auto-assigned
/// - migs!(sql = "SELECT * FROM {{table}}", scope = "...")
/// - migs!(sql = "SELECT * FROM {{table}}", scope = "...", order = 1)
struct MigsInput {
kind: MigsKind,
scope: String,
order: Option<u32>,
}
enum MigsKind {
Path(String),
Sql(String),
}
impl Parse for MigsInput {
fn parse(input: ParseStream) -> syn::Result<Self> {
let lookahead = input.lookahead1();
let kind = if lookahead.peek(syn::Ident) {
let ident: syn::Ident = input.parse()?;
let ident_str = ident.to_string();
if ident_str == "path" {
input.parse::<Token![=]>()?;
let path_lit: LitStr = input.parse()?;
MigsKind::Path(path_lit.value())
} else if ident_str == "sql" {
input.parse::<Token![=]>()?;
let sql_lit: LitStr = input.parse()?;
MigsKind::Sql(sql_lit.value())
} else {
return Err(syn::Error::new(ident.span(), "expected 'path' or 'sql'"));
}
} else {
return Err(lookahead.error());
};
// Parse optional comma and scope, defaulting to "init"
let mut scope = "init".to_string();
let mut order = None;
while input.peek(Token![,]) {
input.parse::<Token![,]>()?;
let ident: syn::Ident = input.parse()?;
let ident_str = ident.to_string();
if ident_str == "scope" {
input.parse::<Token![=]>()?;
let scope_lit: LitStr = input.parse()?;
scope = scope_lit.value();
} else if ident_str == "order" {
input.parse::<Token![=]>()?;
let order_lit: LitInt = input.parse()?;
order = Some(order_lit.base10_parse()?);
} else {
return Err(syn::Error::new(ident.span(), "expected 'scope' or 'order'"));
}
}
Ok(MigsInput { kind, scope, order })
}
}
/// Global registry to track used orders within a single proc macro invocation batch
/// Uses a file-based approach to persist across macro invocations in the same crate
static ORDER_REGISTRY: Mutex<Option<HashSet<(String, u32)>>> = Mutex::new(None);
/// Checks if an order is already used for the given scope
/// Uses a file in the target directory to persist across macro invocations
fn check_duplicate_order(scope: &str, order: u32) -> Result<(), String> {
let mut registry = ORDER_REGISTRY.lock().unwrap();
if registry.is_none() {
// Start with empty registry for each compilation session
// Don't persist across builds to avoid stale data issues
*registry = Some(HashSet::new());
}
let reg = registry.as_mut().unwrap();
let key = (scope.to_string(), order);
if reg.contains(&key) {
return Err(format!(
"Duplicate order '{}' in scope '{}' - another macro already uses this order",
order, scope
));
}
// Add to registry
reg.insert(key);
Ok(())
}
/// The migs! macro registers a SQL script at compile time
/// Usage:
/// - migs!(path = "migrations/001_init.sql") // scope defaults to "init"
/// - migs!(path = "migrations/001_init.sql", scope = "init")
/// - migs!(path = "migrations/001_init.sql", scope = "init", order = 1)
/// - migs!(sql = "CREATE TABLE {{name}} (id INT)") // scope defaults to "init"
/// - migs!(sql = "CREATE TABLE {{name}} (id INT)", scope = "create")
/// - migs!(sql = "CREATE TABLE {{name}} (id INT)", scope = "create", order = 1)
///
/// This macro can be used at module level to register migrations.
/// If an order is specified and another macro already uses that order in the same scope,
/// a compile-time error will be generated.
#[proc_macro]
pub fn migs(input: TokenStream) -> TokenStream {
let input = syn::parse_macro_input!(input as MigsInput);
let (content, source) = match &input.kind {
MigsKind::Path(path) => {
// Read file at compile time
let content = match std::fs::read_to_string(path) {
Ok(c) => c,
Err(e) => {
return syn::Error::new(
Span::call_site(),
format!("Failed to read file '{}': {}", path, e),
)
.to_compile_error()
.into();
}
};
(content, path.clone())
}
MigsKind::Sql(sql) => (sql.clone(), "<inline>".to_string()),
};
let scope = input.scope;
// Check for duplicate order if one was specified
if let Some(order) = input.order {
if let Err(e) = check_duplicate_order(&scope, order) {
return syn::Error::new(Span::call_site(), e)
.to_compile_error()
.into();
}
}
// Generate the order tokens explicitly to handle Option<u32> correctly
let order_tokens = match input.order {
Some(o) => quote! { Some(#o) },
None => quote! { None },
};
// Generate code that registers this migration using inventory
// We construct the Migs directly - strings become &'static str baked into binary
let expanded = quote! {
::migs::inventory::submit! {
::migs::Migs {
scope: #scope,
source: #source,
content: #content,
order: #order_tokens,
}
}
};
TokenStream::from(expanded)
}
/// The collect! macro returns a Vec<&'static Migs> with all registered scripts
#[proc_macro]
pub fn collect(_input: TokenStream) -> TokenStream {
// Use the inventory crate to collect all registered Migs
let expanded = quote! {
{
use ::migs::inventory;
inventory::iter::<::migs::Migs>
.into_iter()
.collect::<Vec<_>>()
}
};
TokenStream::from(expanded)
}