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
14fn 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 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 Meta::Path(_) => Vec::new(),
52 _ => vec![meta.clone()],
54 }
55}
56
57fn 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 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 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 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 let existing_derives: Vec<_> = input_ref
130 .attrs
131 .iter()
132 .filter(|attr| attr.path().is_ident("derive"))
133 .cloned()
134 .collect();
135
136 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 let excluded_keys = collect_excluded_keys(fields, &env_args);
218
219 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 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 pub fn new() -> Self {
261 Self {
262 #params_new_field,
263 #(#builder_field_assigns),*,
264 }
265 }
266 pub fn build(&mut self) -> Result<#struct_name, better_config::Error> {
268 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 trait #helper_trait {
280 #(#getter_methods)*
281 }
282
283 impl better_config::AbstractConfig<#params_type> for #struct_name {
285 fn load(target: Option<String>) -> Result<#params_type, better_config::Error> {
286 <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 <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 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.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 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 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 if let Some(ref prefix) = env_args.prefix {
403 var_name = format!("{}{}", prefix, var_name);
404 }
405
406 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
451fn 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
509fn 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 if let Ok(args) =
529 meta_list.parse_args_with(Punctuated::<Meta, Token![,]>::parse_terminated)
530 {
531 for meta in args {
532 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
546fn collect_excluded_keys(fields: &Fields, env_args: &StructEnvArgs) -> Vec<String> {
556 fields
557 .iter()
558 .filter_map(|field| {
559 if field.attrs.iter().any(|attr| attr.path().is_ident("env")) {
561 return None;
562 }
563
564 if has_no_env_override(field) {
565 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 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}