Skip to main content

better_config_derive/
lib.rs

1use proc_macro::TokenStream;
2use quote::quote;
3use syn::{
4    parse_macro_input, punctuated::Punctuated, Data, DeriveInput, Field, Fields, Meta, Path, Token,
5};
6
7struct StructEnvArgs {
8    trait_path: Option<Path>,
9    prefix: Option<String>,
10    target: Option<String>,
11    generic_args: Vec<syn::GenericArgument>,
12}
13
14/**
15 * `env` macro for deriving environment variable loading functionality.
16 *
17 * ## Example
18 * ```rust,ignore
19 *  #[env(EnvConfig)]
20 *  pub struct Config {
21 *       #[conf(from = "PORT", default = "5432", setter = "set_port", getter = "get_port")]
22 *       pub port: u16,
23 *   }
24 *  ```
25 *
26 */
27fn parse_field_env_args(field: &Field, meta: &Meta) -> Vec<Meta> {
28    if field.attrs.iter().any(|attr| attr.path().is_ident("env")) {
29        return Vec::new();
30    }
31    match meta {
32        // #[env(name = "value", other = "value2")]
33        Meta::List(l) => l
34            .parse_args_with(Punctuated::<Meta, Token![,]>::parse_terminated)
35            .unwrap_or_else(|err| {
36                let field_name = field
37                    .ident
38                    .as_ref()
39                    .map(|i| i.to_string())
40                    .unwrap_or_else(|| String::from("unnamed"));
41
42                panic!(
43                    "Failed to parse env attribute on field `{}`: {:?}",
44                    field_name, err
45                )
46            })
47            .iter()
48            .cloned()
49            .collect(),
50        // #[env]
51        Meta::Path(_) => Vec::new(),
52        // #[env = "value"]
53        _ => vec![meta.clone()],
54    }
55}
56
57// #[env(EnvConfig(prefix = "APP_", target = ".env"))]
58fn parse_struct_env_args(args: Meta) -> StructEnvArgs {
59    let mut prefix = None;
60    let mut target = None;
61    let mut generic_args = Vec::new();
62    let trait_path;
63
64    match args {
65        // #[env(EnvConfig(prefix = "...", ...))]
66        Meta::List(meta_list) => {
67            trait_path = Some(meta_list.path.clone());
68
69            if let Some(last_segment) = meta_list.path.segments.last() {
70                if let syn::PathArguments::AngleBracketed(angle_bracketed) = &last_segment.arguments
71                {
72                    generic_args.extend(angle_bracketed.args.iter().cloned());
73                }
74            }
75
76            let _ = meta_list.parse_nested_meta(|nested_meta| {
77                if nested_meta.path.is_ident("prefix") {
78                    if let Ok(value) = nested_meta.value()?.parse::<syn::LitStr>() {
79                        prefix = Some(value.value());
80                    }
81                } else if nested_meta.path.is_ident("target") {
82                    if let Ok(value) = nested_meta.value()?.parse::<syn::LitStr>() {
83                        target = Some(value.value());
84                    }
85                }
86                Ok(())
87            });
88        }
89        // #[env(EnvConfig)]
90        Meta::Path(path) => {
91            trait_path = Some(path.clone());
92
93            if let Some(last_segment) = path.segments.last() {
94                if let syn::PathArguments::AngleBracketed(angle_bracketed) = &last_segment.arguments
95                {
96                    generic_args.extend(angle_bracketed.args.iter().cloned());
97                }
98            }
99        }
100        _ => panic!(
101            "Invalid env macro arguments. Expected #[env(EnvConfig)] or #[env(EnvConfig(...))]"
102        ),
103    }
104
105    // generic_args > 1  => panic!
106    if generic_args.len() > 1 {
107        panic!("env macro only supports one generic argument");
108    }
109
110    StructEnvArgs {
111        trait_path,
112        prefix,
113        target,
114        generic_args,
115    }
116}
117#[proc_macro_attribute]
118pub fn env(args: TokenStream, input: TokenStream) -> TokenStream {
119    let meta = parse_macro_input!(args as Meta);
120    let env_args = parse_struct_env_args(meta);
121
122    let input_clone = input.clone();
123    let input_ref = parse_macro_input!(input_clone as DeriveInput);
124    let struct_name = &input_ref.ident;
125    let vis = &input_ref.vis;
126    let trait_path = &env_args.trait_path;
127
128    // Extract existing derives
129    let existing_derives: Vec<_> = input_ref
130        .attrs
131        .iter()
132        .filter(|attr| attr.path().is_ident("derive"))
133        .cloned()
134        .collect();
135
136    // Extract fields from the input
137    let fields = match &input_ref.data {
138        Data::Struct(data_struct) => &data_struct.fields,
139        _ => panic!("env macro only supports structs"),
140    };
141
142    let builder_field_assigns = fields
143        .iter()
144        .map(|field| handle_builder_field_assign(&env_args, field));
145
146    let field_defs = fields.iter().map(|field| {
147        let field_name = &field.ident;
148        let field_type = &field.ty;
149        let field_vis = &field.vis;
150        quote! {
151            #field_vis #field_name: #field_type
152        }
153    });
154
155    let builder_field_defs = fields.iter().map(|field| {
156        let field_name = &field.ident;
157        let field_type = &field.ty;
158        let field_vis = &field.vis;
159        quote! {
160            #field_vis #field_name: Option<#field_type>
161        }
162    });
163
164    let target = match &env_args.target {
165        Some(t) => quote! { Some(#t.to_string()) },
166        None => quote! { None },
167    };
168
169    let params_type = if env_args.generic_args.is_empty() {
170        quote! { ::std::collections::HashMap<String, String> }
171    } else {
172        let generic_arg = &env_args.generic_args[0];
173        quote! { #generic_arg }
174    };
175
176    let params_field = quote! {
177        _params: ::std::collections::HashMap<String, String>
178    };
179
180    let params_new_field = quote! {
181        _params: ::std::collections::HashMap::new()
182    };
183
184    let helper_trait = quote::format_ident!("{}BetterHelper", struct_name);
185
186    let struct_builder = quote::format_ident!("{}Builder", struct_name);
187
188    let loaded_params_var = quote::format_ident!("loaded_params");
189    let field_assigns = handle_field_assigns(fields, &env_args, &loaded_params_var);
190
191    let getter_methods = fields.iter().filter_map(|field| {
192        let field_env_attr = field.attrs.iter().find(|attr| attr.path().is_ident("conf"));
193        if let Some(field_env_attr) = field_env_attr {
194            let field_env_args = parse_field_env_args(field, &field_env_attr.meta);
195            for attr in field_env_args {
196                if attr.path().is_ident("getter") {
197                    if let Meta::NameValue(name_value) = attr {
198                        if let syn::Expr::Lit(syn::ExprLit {
199                            lit: syn::Lit::Str(lit_str),
200                            ..
201                        }) = &name_value.value
202                        {
203                            let getter_ident = quote::format_ident!("{}", lit_str.value());
204                            let field_type = &field.ty;
205                            return Some(quote! {
206                                fn #getter_ident(&self,#params_field) -> #field_type;
207                            });
208                        }
209                    }
210                }
211            }
212        }
213        None
214    });
215
216    // Collect excluded keys for no_env_override
217    let excluded_keys = collect_excluded_keys(fields, &env_args);
218
219    // Generate the load call - use load_with_override if there are excluded keys
220    // and the trait supports it (file-based loaders)
221    let load_call = if excluded_keys.is_empty() {
222        quote! {
223            <Self as #trait_path<#params_type>>::load(#target)?
224        }
225    } else {
226        let excluded_keys_tokens: Vec<_> = excluded_keys
227            .iter()
228            .map(|k| quote! { #k.to_string() })
229            .collect();
230        quote! {
231            {
232                let mut excluded = ::std::collections::HashSet::new();
233                #(excluded.insert(#excluded_keys_tokens);)*
234                <Self as #trait_path<#params_type>>::load_with_override(#target, &excluded)?
235            }
236        }
237    };
238
239    let expanded = quote! {
240        #(#existing_derives)*
241        #vis struct #struct_name {
242            #params_field,
243            #(#field_defs),*,
244        }
245
246        impl #struct_name {
247            // builder
248            pub fn builder() -> #struct_builder {
249                #struct_builder::new()
250            }
251        }
252
253        #vis struct #struct_builder {
254            #params_field,
255            #(#builder_field_defs),*,
256        }
257
258        impl #struct_builder {
259            // new
260            pub fn new() -> Self {
261                Self {
262                    #params_new_field,
263                    #(#builder_field_assigns),*,
264                }
265            }
266            // builder methods
267            pub fn build(&mut self) -> Result<#struct_name, better_config::Error> {
268                // load first (with excluded keys if any)
269                let loaded_params = #load_call;
270                let config = #struct_name {
271                    _params: loaded_params.clone(),
272                    #(#field_assigns),*,
273                };
274                Ok(config)
275            }
276        }
277
278        // generate a help trait to add getter and setter methods
279        trait #helper_trait {
280            #(#getter_methods)*
281        }
282
283        // First implement AbstractConfig
284        impl better_config::AbstractConfig<#params_type> for #struct_name {
285            fn load(target: Option<String>) -> Result<#params_type, better_config::Error> {
286                // Default to calling EnvConfig's load
287                <Self as #trait_path<#params_type>>::load(#target)
288            }
289        }
290
291        impl better_config::AbstractConfig<#params_type> for #struct_builder {
292            fn load(target: Option<String>) -> Result<#params_type, better_config::Error> {
293                // Default to calling EnvConfig's load
294                <Self as #trait_path<#params_type>>::load(#target)
295            }
296        }
297
298        impl #trait_path<#params_type> for #struct_name  {}
299        impl #trait_path<#params_type> for #struct_builder  {}
300
301    };
302
303    TokenStream::from(expanded)
304}
305
306fn handle_builder_field_assign(
307    env_args: &StructEnvArgs,
308    field: &Field,
309) -> proc_macro2::TokenStream {
310    let field_name = &field.ident;
311    let field_type = &field.ty;
312
313    // nested
314    let is_nested = field.attrs.iter().any(|attr| attr.path().is_ident("env"));
315    if is_nested {
316        return quote! {
317            #field_name: None
318        };
319    }
320
321    let from = get_var_name(field, "from");
322    let default = get_var_name(field, "default");
323    // if from and default are both None, return None
324    if from.is_none() && default.is_none() {
325        return quote! {
326            #field_name: None
327        };
328    }
329
330    let mut var_name =
331        from.unwrap_or_else(|| field_name.as_ref().unwrap().to_string().to_uppercase());
332
333    // Add prefix if specified
334    if let Some(ref prefix) = env_args.prefix {
335        var_name = format!("{}{}", prefix, var_name);
336    }
337
338    if let Some(default) = default {
339        return quote! {
340            #field_name: ::better_config::utils::env::get_optional_or::<_,#field_type>(#var_name, #default.parse::<#field_type>().unwrap())
341        };
342    }
343
344    quote! {
345            #field_name: ::better_config::utils::env::get_optional::<_,#field_type>(#var_name)
346    }
347}
348
349fn handle_field_assign(
350    env_args: &StructEnvArgs,
351    field: &Field,
352    loaded_params_var: &proc_macro2::Ident,
353) -> proc_macro2::TokenStream {
354    // get field conf attr
355    let field_env_attr = field.attrs.iter().find(|attr| attr.path().is_ident("conf"));
356
357    let field_name = &field.ident;
358
359    let is_nested = field.attrs.iter().any(|attr| attr.path().is_ident("env"));
360
361    if is_nested {
362        let field_type = &field.ty;
363        return quote! {
364            #field_name: #field_type::builder()
365                .build()
366                .expect("Failed to build nested config")
367        };
368    }
369
370    let assign = if let Some(field_env_attr) = field_env_attr {
371        match &field_env_attr.meta {
372            Meta::List(_) => handle_field_meta_list(env_args, field, loaded_params_var),
373            _ => panic!(
374                "Unsupported env attribute on field `{}`",
375                field_name.as_ref().unwrap()
376            ),
377        }
378    } else {
379        let field_name_str = field_name.as_ref().unwrap().to_string().to_uppercase();
380        quote! {
381            #field_name: ::better_config::utils::env::get_or_else(#field_name_str, || panic!("Failed to load from var: {}", #field_name_str))?
382        }
383    };
384
385    quote! {
386        #assign
387    }
388}
389
390fn handle_field_meta_list(
391    env_args: &StructEnvArgs,
392    field: &Field,
393    loaded_params_var: &proc_macro2::Ident,
394) -> proc_macro2::TokenStream {
395    let field_name = &field.ident;
396    let field_type = &field.ty;
397
398    let mut var_name = get_var_name(field, "from")
399        .unwrap_or_else(|| field_name.as_ref().unwrap().to_string().to_uppercase());
400
401    // Add prefix if specified
402    if let Some(ref prefix) = env_args.prefix {
403        var_name = format!("{}{}", prefix, var_name);
404    }
405
406    // handle attributes
407    let default = get_var_name(field, "default");
408    let setter_name = get_var_name(field, "setter");
409    let getter_name = get_var_name(field, "getter");
410
411    if let Some(getter) = getter_name {
412        let getter_ident = quote::format_ident!("{}", getter);
413        quote! {
414            #field_name: <Self>::#getter_ident(&self,&#loaded_params_var)
415        }
416    } else if let Some(setter) = setter_name {
417        let setter_ident = quote::format_ident!("{}", setter);
418        quote! {
419            #field_name: {
420                let value = #loaded_params_var.get(#var_name).cloned().unwrap_or_default();
421                self.#setter_ident(value.clone());
422                value
423            }
424        }
425    } else if let Some(default) = default {
426        quote! {
427            #field_name: #loaded_params_var.get(#var_name)
428                .and_then(|v| v.parse::<#field_type>().ok())
429                .unwrap_or_else(|| #default.parse::<#field_type>().unwrap())
430
431        }
432    } else {
433        quote! {
434            #field_name: #loaded_params_var.get(#var_name)
435                .and_then(|v| v.parse::<#field_type>().ok())
436                .unwrap_or_default()
437        }
438    }
439}
440
441fn handle_field_assigns<'a>(
442    fields: &'a Fields,
443    env_args: &'a StructEnvArgs,
444    loaded_params_var: &'a proc_macro2::Ident,
445) -> impl Iterator<Item = proc_macro2::TokenStream> + 'a {
446    fields
447        .iter()
448        .map(move |field| handle_field_assign(env_args, field, loaded_params_var))
449}
450
451/// Extracts the variable name from the field attributes, looking for a specific attribute name.
452/// If the attribute is not found, it returns `None`.
453///
454/// # Arguments
455/// * `field` - The field from which to extract the variable name.
456/// * `field_name` - The name of the attribute to look for.
457/// # Returns
458/// * `Option<String>` - The variable name if found, otherwise `None`.
459///
460/// # Example
461/// ```rust,ignore
462/// let field = ...; // Some syn::Field
463/// let var_name = get_var_name(&field, "from");
464/// if let Some(name) = var_name {
465///     println!("Found variable name: {}", name);
466/// } else {
467///     println!("Variable name not found.");
468/// }
469/// ```
470fn get_var_name(field: &Field, field_name: &'static str) -> Option<String> {
471    for attr in &field.attrs {
472        if attr.path().is_ident("conf") {
473            if let Meta::List(meta_list) = &attr.meta {
474                if let Ok(args) =
475                    meta_list.parse_args_with(Punctuated::<Meta, Token![,]>::parse_terminated)
476                {
477                    for meta in args {
478                        if let Meta::NameValue(name_value) = meta {
479                            if name_value.path.is_ident(field_name) {
480                                if let syn::Expr::Lit(syn::ExprLit {
481                                    lit: syn::Lit::Str(lit_str),
482                                    ..
483                                }) = &name_value.value
484                                {
485                                    return Some(lit_str.value());
486                                }
487                            }
488                        }
489                    }
490                }
491                let mut result = None;
492                let _ = meta_list.parse_nested_meta(|meta| {
493                    if meta.path.is_ident(field_name) {
494                        if let Ok(value) = meta.value()?.parse::<syn::LitStr>() {
495                            result = Some(value.value());
496                        }
497                    }
498                    Ok(())
499                });
500                if result.is_some() {
501                    return result;
502                }
503            }
504        }
505    }
506    None
507}
508
509/// Checks if a field has the `no_env_override` attribute set.
510///
511/// # Arguments
512/// * `field` - The field to check for the `no_env_override` attribute.
513///
514/// # Returns
515/// * `bool` - `true` if the field has `no_env_override`, `false` otherwise.
516///
517/// # Example
518/// ```rust,ignore
519/// // For a field with #[conf(from = "KEY", no_env_override)]
520/// let has_no_override = has_no_env_override(&field);
521/// assert!(has_no_override);
522/// ```
523fn has_no_env_override(field: &Field) -> bool {
524    for attr in &field.attrs {
525        if attr.path().is_ident("conf") {
526            if let Meta::List(meta_list) = &attr.meta {
527                // Try parsing as comma-separated meta items
528                if let Ok(args) =
529                    meta_list.parse_args_with(Punctuated::<Meta, Token![,]>::parse_terminated)
530                {
531                    for meta in args {
532                        // Check for #[conf(no_env_override)] - path-style attribute
533                        if let Meta::Path(path) = &meta {
534                            if path.is_ident("no_env_override") {
535                                return true;
536                            }
537                        }
538                    }
539                }
540            }
541        }
542    }
543    false
544}
545
546/// Collects all field keys that have the `no_env_override` attribute.
547/// These keys should be excluded from environment variable override.
548///
549/// # Arguments
550/// * `fields` - The fields to check.
551/// * `env_args` - The struct-level env arguments (for prefix handling).
552///
553/// # Returns
554/// * `Vec<String>` - List of keys that should not be overridden by env vars.
555fn collect_excluded_keys(fields: &Fields, env_args: &StructEnvArgs) -> Vec<String> {
556    fields
557        .iter()
558        .filter_map(|field| {
559            // Skip nested fields
560            if field.attrs.iter().any(|attr| attr.path().is_ident("env")) {
561                return None;
562            }
563
564            if has_no_env_override(field) {
565                // Get the key name (from attribute or field name)
566                let key = get_var_name(field, "from").unwrap_or_else(|| {
567                    field
568                        .ident
569                        .as_ref()
570                        .map(|i| i.to_string().to_uppercase())
571                        .unwrap_or_default()
572                });
573
574                // Apply prefix if specified (for consistency with how keys are stored)
575                let full_key = if let Some(ref prefix) = env_args.prefix {
576                    format!("{}{}", prefix, key)
577                } else {
578                    key
579                };
580
581                Some(full_key)
582            } else {
583                None
584            }
585        })
586        .collect()
587}