prometheus-derive-macros 0.1.0

Procedural macros for prometheus with automatic metric registration
Documentation
use proc_macro2::Span;
use quote::quote;
use syn::Ident;

use crate::parsing::MetricDefinition;

/// Entrypoint of all the code generation behind the function-like macro, called
/// after the full parsing of the macro
pub fn generate_metric_code(metric: &MetricDefinition) -> proc_macro2::TokenStream {
    let metric_ident = Ident::new(&metric.variable_name, Span::call_site());
    let help_str = &metric.help;
    let metric_name = &metric.name;

    if metric.labels.is_empty() {
        return generate_metric_var(&metric_ident, &metric.metric_type, metric_name, help_str);
    }

    generate_metric_var_with_labels(
        &metric_ident,
        &metric.metric_type,
        &metric.labels,
        metric_name,
        help_str,
    )
}

/// Generates the static variable for a metric that doesn't have any label
///
/// This kind of metric shouldn't be encapsulated into a MetricFamily
fn generate_metric_var(
    metric_ident: &Ident,
    metric_type: &str,
    metric_name: &str,
    help_str: &str,
) -> proc_macro2::TokenStream {
    let metric_type_path = get_metric_type_path(metric_type);

    quote! {
        pub static #metric_ident: LazyLock<#metric_type_path> = LazyLock::new(|| {
            let metric = #metric_type_path::default();
            // Since lazy loaded, we need to register it to the global registry
            // when first accessed
            if let Ok(mut registry) = prometheus_derive::GLOBAL_REGISTRY.write() {
                registry.register(#metric_name, #help_str, metric.clone());
            }
            metric
        });
    }
}

/// Generates the static variable for a metric that has labels, allowing
/// to group the cardinality of the metric
fn generate_metric_var_with_labels(
    metric_ident: &Ident,
    metric_type: &str,
    labels: &[(String, String)],
    metric_name: &str,
    help_str: &str,
) -> proc_macro2::TokenStream {
    let struct_name = get_labels_struct_name(metric_ident);

    let label_struct = generate_label_struct(metric_ident, labels);
    let metric_type_path = get_metric_type_path(metric_type);

    let family_construction = quote! {
        let family = Family::<#struct_name, #metric_type_path>::default();
    };

    quote! {
        #label_struct

        pub static #metric_ident: LazyLock<Family<#struct_name, #metric_type_path>> = LazyLock::new(|| {
            #family_construction

            // Since lazy loaded, we need to register it to the global registry
            // when first accessed
            if let Ok(mut registry) = prometheus_derive::GLOBAL_REGISTRY.write() {
                registry.register(#metric_name, #help_str, family.clone());
            }
            family
        });
    }
}

fn get_labels_struct_name(metric_ident: &Ident) -> Ident {
    let struct_name_str = to_pascal_case(&metric_ident.to_string()) + "Labels";
    Ident::new(&struct_name_str, Span::call_site())
}

fn generate_label_struct(
    metric_ident: &Ident,
    labels: &[(String, String)],
) -> proc_macro2::TokenStream {
    if labels.is_empty() {
        panic!("It shouldn't have tried to generate an empty label struct, please report this as a bug");
    }

    let struct_name = get_labels_struct_name(metric_ident);
    let fields = generate_struct_fields(labels);
    let constructor = generate_constructor(&struct_name, labels);

    quote! {
        #[derive(Clone, Debug, Hash, PartialEq, Eq, prometheus_client::encoding::EncodeLabelSet)]
        pub struct #struct_name {
            #(#fields),*
        }

        #constructor
    }
}

fn generate_struct_fields(labels: &[(String, String)]) -> Vec<proc_macro2::TokenStream> {
    labels
        .iter()
        .map(|(name, type_str)| {
            let field_name = Ident::new(name, Span::call_site());
            let field_type = parse_type_string(type_str);
            quote! { pub #field_name: #field_type }
        })
        .collect()
}

fn generate_constructor(
    struct_name: &Ident,
    labels: &[(String, String)],
) -> proc_macro2::TokenStream {
    let constructor_params = generate_constructor_params(labels);
    let constructor_fields = generate_constructor_fields(labels);

    quote! {
        impl #struct_name {
            pub fn new(#(#constructor_params),*) -> Self {
                Self {
                    #(#constructor_fields),*
                }
            }
        }
    }
}

fn generate_constructor_params(labels: &[(String, String)]) -> Vec<proc_macro2::TokenStream> {
    labels
        .iter()
        .map(|(name, type_str)| {
            let param_name = Ident::new(name, Span::call_site());
            let param_type = parse_type_string(type_str);
            quote! { #param_name: #param_type }
        })
        .collect()
}

fn generate_constructor_fields(labels: &[(String, String)]) -> Vec<proc_macro2::TokenStream> {
    labels
        .iter()
        .map(|(name, _)| {
            let field_name = Ident::new(name, Span::call_site());
            quote! { #field_name }
        })
        .collect()
}

/// Get the appropriate metric type path for code generation.
/// Uses the type as provided by the user (they should have imported it).
fn get_metric_type_path(metric_type: &str) -> proc_macro2::TokenStream {
    let metric_type_ident = Ident::new(metric_type, Span::call_site());
    quote! { #metric_type_ident }
}

/// Convert snake_case or SCREAMING_SNAKE_CASE to PascalCase.
pub fn to_pascal_case(snake_case: &str) -> String {
    snake_case
        .to_lowercase()
        .split('_')
        .map(|word| {
            let mut chars = word.chars();
            match chars.next() {
                None => String::new(),
                Some(first) => first.to_uppercase().collect::<String>() + chars.as_str(),
            }
        })
        .collect()
}

/// Parse a type string into a TokenStream for code generation.
pub fn parse_type_string(type_str: &str) -> proc_macro2::TokenStream {
    match type_str {
        "String" => quote! { String },
        "u16" => quote! { u16 },
        "u32" => quote! { u32 },
        "u64" => quote! { u64 },
        "i16" => quote! { i16 },
        "i32" => quote! { i32 },
        "i64" => quote! { i64 },
        "f32" => quote! { f32 },
        "f64" => quote! { f64 },
        "bool" => quote! { bool },
        _ => {
            // For custom types, try to parse as TokenStream
            type_str.parse().unwrap_or_else(|_| quote! { String })
        }
    }
}

#[cfg(test)]
mod tests {
    use super::*;
    use crate::parsing::MetricDefinition;

    fn create_test_metric(
        name: &str,
        metric_type: &str,
        labels: Vec<(String, String)>,
    ) -> MetricDefinition {
        MetricDefinition {
            variable_name: name.to_string(),
            metric_type: metric_type.to_string(),
            name: name.to_string(),
            help: format!("Test metric: {}", name),
            labels,
        }
    }

    #[test]
    fn test_to_pascal_case() {
        assert_eq!(to_pascal_case("snake_case"), "SnakeCase");
        assert_eq!(to_pascal_case("http_requests_total"), "HttpRequestsTotal");
        assert_eq!(to_pascal_case("HTTP_REQUESTS_TOTAL"), "HttpRequestsTotal");
        assert_eq!(to_pascal_case("single"), "Single");
        assert_eq!(to_pascal_case("SINGLE"), "Single");
        assert_eq!(to_pascal_case(""), "");
    }

    #[test]
    fn test_parse_type_string() {
        let string_type = parse_type_string("String");
        assert_eq!(string_type.to_string(), "String");

        let u32_type = parse_type_string("u32");
        assert_eq!(u32_type.to_string(), "u32");
    }
    #[test]
    fn test_generate_simple_metric_code() {
        let metric = create_test_metric("simple_counter", "Counter", vec![]);
        let code = generate_metric_code(&metric);
        let code_str = code.to_string();

        assert!(code_str.contains("simple_counter"));
        assert!(code_str.contains("Counter"));
        assert!(code_str.contains("LazyLock"));
    }

    #[test]
    fn test_generate_labeled_metric_code() {
        let labels = vec![
            ("method".to_string(), "String".to_string()),
            ("status".to_string(), "u16".to_string()),
        ];
        let metric = create_test_metric("http_requests", "Counter", labels);
        let code = generate_metric_code(&metric);
        let code_str = code.to_string();

        assert!(code_str.contains("http_requests"));
        assert!(code_str.contains("HttpRequestsLabels"));
        assert!(code_str.contains("Family"));
        assert!(code_str.contains("Counter"));
        assert!(code_str.contains("EncodeLabelSet"));
    }

    #[test]
    #[should_panic]
    fn test_generate_label_struct_empty() {
        let struct_name = Ident::new("TestLabels", Span::call_site());
        generate_label_struct(&struct_name, &[]);
    }

    #[test]
    fn test_generate_label_struct_with_fields() {
        let struct_name = Ident::new("Http", Span::call_site());
        let labels = vec![
            ("method".to_string(), "String".to_string()),
            ("status".to_string(), "u16".to_string()),
        ];
        let code = generate_label_struct(&struct_name, &labels);
        let code_str = code.to_string();

        assert!(code_str.contains("struct HttpLabels"));
        assert!(code_str.contains("pub method : String"));
        assert!(code_str.contains("pub status : u16"));
        assert!(code_str.contains("pub fn new"));
    }

    #[test]
    fn test_get_metric_type_path_counter() {
        let path = get_metric_type_path("Counter");
        assert_eq!(path.to_string(), "Counter");
    }

    #[test]
    fn test_get_metric_type_path_gauge() {
        let path = get_metric_type_path("Gauge");
        assert_eq!(path.to_string(), "Gauge");
    }

    #[test]
    fn test_get_metric_type_path_custom() {
        let path = get_metric_type_path("CustomMetric");
        assert_eq!(path.to_string(), "CustomMetric");
    }

    #[test]
    fn test_generate_constructor_params() {
        let labels = vec![
            ("method".to_string(), "String".to_string()),
            ("timeout".to_string(), "u64".to_string()),
        ];
        let params = generate_constructor_params(&labels);

        assert_eq!(params.len(), 2);
        assert!(params[0].to_string().contains("method : String"));
        assert!(params[1].to_string().contains("timeout : u64"));
    }

    #[test]
    fn test_generate_constructor_fields() {
        let labels = vec![
            ("env".to_string(), "String".to_string()),
            ("region".to_string(), "String".to_string()),
        ];
        let fields = generate_constructor_fields(&labels);

        assert_eq!(fields.len(), 2);
        assert_eq!(fields[0].to_string(), "env");
        assert_eq!(fields[1].to_string(), "region");
    }

    #[test]
    fn test_generate_struct_fields() {
        let labels = vec![
            ("name".to_string(), "String".to_string()),
            ("count".to_string(), "u32".to_string()),
        ];
        let fields = generate_struct_fields(&labels);

        assert_eq!(fields.len(), 2);
        assert!(fields[0].to_string().contains("pub name : String"));
        assert!(fields[1].to_string().contains("pub count : u32"));
    }
}