Skip to content
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)
}