diff --git a/derive/src/action.rs b/derive/src/action.rs index f83a502..4fea590 100644 --- a/derive/src/action.rs +++ b/derive/src/action.rs @@ -15,14 +15,14 @@ pub(crate) enum ActionType { Map(Vec), } -fn parse_paths(attr: Attribute) -> Vec { +fn parse_paths(attr: &Attribute) -> Vec { attr.parse_args_with(Punctuated::::parse_terminated) .into_iter() .flatten() .collect() } -pub(crate) fn parse_action_attr(attr: Attribute) -> Option { +pub(crate) fn parse_action_attr(attr: &Attribute) -> Option { if attr.path.is_ident("collect") { let inner: ActionType = attr.parse_args().unwrap(); Some(ActionAttr { diff --git a/derive/src/attributes.rs b/derive/src/attributes.rs index 13c61c7..2e02d55 100644 --- a/derive/src/attributes.rs +++ b/derive/src/attributes.rs @@ -30,6 +30,7 @@ enum AttributeArguments { Value(Expr), NumArgs(RangeInclusive), File(String), + Env(String), Last, } @@ -70,6 +71,28 @@ impl OptionAttr { } } +#[derive(Default)] +pub(crate) struct FieldAttr { + pub(crate) default: Option, + pub(crate) env: Option, +} + +impl FieldAttr { + pub(crate) fn parse(attr: &Attribute) -> Self { + let mut field_attr = Self::default(); + + for arg in AttributeArguments::parse_all(attr) { + match arg { + AttributeArguments::Default(e) => field_attr.default = Some(e), + AttributeArguments::Env(e) => field_attr.env = Some(e), + _ => panic!("Invalid argument"), + }; + } + + field_attr + } +} + #[derive(Default)] pub(crate) struct ValueAttr { pub(crate) keys: Vec, @@ -227,6 +250,7 @@ impl Parse for AttributeArguments { "default" => return Ok(Self::Default(input.parse::()?)), "value" => return Ok(Self::Value(input.parse::()?)), "file" => return Ok(Self::File(input.parse::()?.value())), + "env" => return Ok(Self::Env(input.parse::()?.value())), _ => panic!("Unrecognized argument {} for option attribute", name), }; } diff --git a/derive/src/field.rs b/derive/src/field.rs new file mode 100644 index 0000000..6a33489 --- /dev/null +++ b/derive/src/field.rs @@ -0,0 +1,104 @@ +use proc_macro2::TokenStream; +use quote::{quote, ToTokens}; +use syn::{Attribute, Field, Ident}; + +use crate::{ + action::{parse_action_attr, ActionAttr, ActionType}, + attributes::FieldAttr, +}; + +pub(crate) struct FieldData { + pub(crate) ident: Ident, + pub(crate) default_value: TokenStream, + pub(crate) match_stmt: TokenStream, +} + +pub(crate) fn parse_field(field: &Field) -> FieldData { + let field_ident = field.ident.as_ref().unwrap().clone(); + + let field_attr = parse_field_attr(&field.attrs); + + let mut default_value = match field_attr.default { + Some(val) => val.to_token_stream(), + None => quote!(::core::default::Default::default()), + }; + + if let Some(env_var) = field_attr.env { + default_value = quote!( + match ::std::env::var_os(#env_var) { + Some(x) => ::uutils_args::FromValue::from_value("", x)?, + None => #default_value + } + ) + } + + let match_arms = field + .attrs + .iter() + .filter_map(parse_action_attr) + .flat_map(|attr| action_attr_to_match_arms(&field_ident, attr)); + + let match_stmt = quote!(match arg.clone() { + #(#match_arms)*, + _ => {} + }); + + FieldData { + ident: field_ident, + default_value, + match_stmt, + } +} + +pub(crate) fn parse_field_attr(attrs: &[Attribute]) -> FieldAttr { + for attr in attrs { + if attr.path.is_ident("field") { + return FieldAttr::parse(attr); + } + } + FieldAttr::default() +} + +fn action_attr_to_match_arms(field_ident: &Ident, attr: ActionAttr) -> Vec { + let mut match_arms = Vec::new(); + match attr.action_type { + ActionType::Map(arms) => { + for arm in arms { + match_arms.push(field_expression( + arm.pat.to_token_stream(), + arm.body.to_token_stream(), + field_ident, + attr.collect, + )); + } + } + + ActionType::Set(pats) => { + let pats: Vec<_> = pats.iter().map(|p| quote!(#p(x))).collect(); + match_arms.push(field_expression( + quote!(#(#pats)|*), + quote!(x), + field_ident, + attr.collect, + )); + } + }; + match_arms +} + +fn field_expression( + pat: TokenStream, + expr: TokenStream, + field_ident: &Ident, + collect: bool, +) -> TokenStream { + if collect { + quote!( + #pat => { self.#field_ident.push(#expr) } + ) + } else { + quote!( + #pat => { self.#field_ident = #expr } + ) + } +} diff --git a/derive/src/lib.rs b/derive/src/lib.rs index 3cbfcb2..5d9fcfd 100644 --- a/derive/src/lib.rs +++ b/derive/src/lib.rs @@ -1,13 +1,14 @@ mod action; mod argument; mod attributes; +mod field; mod flags; mod help; mod markdown; -use action::{parse_action_attr, ActionAttr, ActionType}; use argument::{long_handling, parse_argument, positional_handling, short_handling}; use attributes::ValueAttr; +use field::{parse_field, FieldData}; use help::{help_handling, help_string, parse_help_attr, parse_version_attr, version_handling}; use proc_macro::TokenStream; @@ -19,7 +20,7 @@ use syn::{ DeriveInput, Fields, }; -#[proc_macro_derive(Options, attributes(arg_type, map, set, set_true, set_false, collect))] +#[proc_macro_derive(Options, attributes(arg_type, map, set, field, collect))] pub fn options(input: TokenStream) -> TokenStream { let input = parse_macro_input!(input as DeriveInput); @@ -44,50 +45,26 @@ pub fn options(input: TokenStream) -> TokenStream { // The key of this map is a literal pattern and the value // is whatever code needs to be run when that pattern is encountered. let mut stmts = Vec::new(); - + let mut defaults = Vec::new(); for field in fields.named { - let field_ident = field.ident.as_ref().unwrap(); - let mut match_arms = vec![]; - for attr in field.attrs { - let Some(ActionAttr { action_type, collect }) = parse_action_attr(attr) else { continue; }; + let FieldData { + ident, + default_value, + match_stmt, + } = parse_field(&field); - let mut patterns_and_expressions = vec![]; - match action_type { - ActionType::Map(arms) => { - for arm in arms { - let pat = arm.pat; - let expr = arm.body; - patterns_and_expressions.push((quote!(#pat), quote!(#expr))); - } - } - - ActionType::Set(pats) => { - let pats: Vec<_> = pats.iter().map(|p| quote!(#p(x))).collect(); - let pats = quote!(#(#pats)|*); - patterns_and_expressions.push((pats, quote!(x.clone()))) - } - }; - for (pat, expr) in patterns_and_expressions { - match_arms.push(if collect { - quote!( - #pat => { self.#field_ident.push(#expr) } - ) - } else { - quote!( - #pat => { self.#field_ident = #expr } - ) - }); - } - } - - stmts.push(quote!(match arg.clone() { - #(#match_arms)* - _ => {} - })) + defaults.push(quote!(#ident: #default_value)); + stmts.push(match_stmt); } let expanded = quote!( impl #impl_generics Options for #name #ty_generics #where_clause { + fn initial() -> Result { + Ok(Self { + #(#defaults),* + }) + } + fn apply_args(&mut self, args: I) -> Result<(), uutils_args::Error> where I: IntoIterator + 'static, diff --git a/src/lib.rs b/src/lib.rs index 2b8f83a..a919f6e 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -119,11 +119,13 @@ pub trait Options: Sized + Default { I: IntoIterator + 'static, I::Item: Into, { - let mut _self = Self::default(); + let mut _self = Self::initial()?; _self.apply_args(args)?; Ok(_self) } + fn initial() -> Result; + fn apply_args(&mut self, args: I) -> Result<(), Error> where I: IntoIterator + 'static, diff --git a/tests/coreutils/base32.rs b/tests/coreutils/base32.rs index 16f03ea..cc1d786 100644 --- a/tests/coreutils/base32.rs +++ b/tests/coreutils/base32.rs @@ -19,7 +19,7 @@ enum Arg { File(PathBuf), } -#[derive(Options)] +#[derive(Options, Default)] #[arg_type(Arg)] struct Settings { #[map(Arg::Decode => true)] @@ -32,23 +32,13 @@ struct Settings { Arg::Wrap(0) => None, Arg::Wrap(n) => Some(n), )] + #[field(default = Some(76))] wrap: Option, #[map(Arg::File(f) => Some(f))] file: Option, } -impl Default for Settings { - fn default() -> Self { - Self { - decode: false, - ignore_garbage: false, - wrap: Some(76), - file: None, - } - } -} - #[test] fn wrap() { assert_eq!(Settings::parse(["base32"]).unwrap().wrap, Some(76)); diff --git a/tests/defaults.rs b/tests/defaults.rs new file mode 100644 index 0000000..4184112 --- /dev/null +++ b/tests/defaults.rs @@ -0,0 +1,47 @@ +use uutils_args::{Arguments, Options}; + +#[test] +fn true_default() { + #[derive(Arguments, Clone)] + enum Arg { + #[option("--foo")] + Foo, + } + + #[derive(Default, Options)] + #[arg_type(Arg)] + struct Settings { + #[map(Arg::Foo => false)] + #[field(default = true)] + foo: bool, + } + + assert!(Settings::parse(["test"]).unwrap().foo); + assert!(!Settings::parse(["test", "--foo"]).unwrap().foo); +} + +#[test] +fn env_var_string() { + #[derive(Arguments, Clone)] + enum Arg { + #[option("--foo=MSG")] + Foo(String), + } + + #[derive(Default, Options)] + #[arg_type(Arg)] + struct Settings { + #[set(Arg::Foo)] + #[field(env = "FOO")] + foo: String, + } + + std::env::set_var("FOO", "one"); + assert_eq!(Settings::parse(["test"]).unwrap().foo, "one"); + + std::env::set_var("FOO", "two"); + assert_eq!(Settings::parse(["test"]).unwrap().foo, "two"); + + std::env::remove_var("FOO"); + assert_eq!(Settings::parse(["test"]).unwrap().foo, ""); +}