From 07c482538e1d09155bf2b7235b975c97b5fe179d Mon Sep 17 00:00:00 2001 From: Terts Diepraam Date: Sun, 27 Nov 2022 22:48:05 +0100 Subject: [PATCH] add options --- derive/src/attributes.rs | 162 ++++++++++++++++++++++++++++++++++ derive/src/lib.rs | 184 ++++++++++++++++++++++----------------- src/lib.rs | 16 ++++ tests/flags.rs | 146 ++++++++++++++++++++++++++++++- 4 files changed, 427 insertions(+), 81 deletions(-) create mode 100644 derive/src/attributes.rs diff --git a/derive/src/attributes.rs b/derive/src/attributes.rs new file mode 100644 index 0000000..ec44950 --- /dev/null +++ b/derive/src/attributes.rs @@ -0,0 +1,162 @@ +use syn::{ + parse::{Parse, ParseStream}, + punctuated::Punctuated, + Attribute, Ident, LitStr, Token, +}; + +use crate::{Arg, Expr}; + +#[derive(Default)] +pub(crate) struct FlagAttr { + pub(crate) flags: Vec, + pub(crate) value: Option, +} + +enum FlagAttrArg { + Arg(Arg), + Value(Expr), +} + +#[derive(Default)] +pub(crate) struct OptionAttr { + pub(crate) flags: Vec, + // This should probably not accept any expr to give better errors. + // Closures should be allowed though. + pub(crate) parser: Option, +} + +enum OptionAttrArg { + Arg(Arg), + Parser(Expr), +} + +#[derive(Default)] +pub(crate) struct ValueAttr { + pub(crate) keys: Vec, + pub(crate) value: Option, +} + +enum ValueAttrArg { + Key(String), + Value(Expr), +} + +pub(crate) fn parse_flag_attr(attr: Attribute) -> FlagAttr { + let mut flag_attr = FlagAttr::default(); + let Ok(parsed_args) = attr + .parse_args_with(Punctuated::::parse_terminated) + else { + return flag_attr; + }; + for arg in parsed_args { + match arg { + FlagAttrArg::Arg(a) => flag_attr.flags.push(a), + FlagAttrArg::Value(e) => flag_attr.value = Some(e), + }; + } + flag_attr +} + +impl Parse for FlagAttrArg { + fn parse(input: ParseStream) -> syn::Result { + if input.peek(LitStr) { + return parse_flag(input).map(Self::Arg); + } + + if input.peek(Ident) { + let name = input.parse::()?.to_string(); + input.parse::()?; + match name.as_str() { + "value" => return Ok(Self::Value(input.parse::()?)), + _ => panic!("Unrecognized argument {} for flag attribute", name), + }; + } + panic!("Arguments to flag attribute must be string literals"); + } +} + +pub(crate) fn parse_option_attr(attr: Attribute) -> OptionAttr { + let mut option_attr = OptionAttr::default(); + let Ok(parsed_args) = attr + .parse_args_with(Punctuated::::parse_terminated) + else { + return option_attr; + }; + + for arg in parsed_args { + match arg { + OptionAttrArg::Arg(a) => option_attr.flags.push(a), + OptionAttrArg::Parser(e) => option_attr.parser = Some(e), + }; + } + option_attr +} + +impl Parse for OptionAttrArg { + fn parse(input: ParseStream) -> syn::Result { + if input.peek(LitStr) { + return parse_flag(input).map(Self::Arg); + } + + if input.peek(Ident) { + let name = input.parse::()?.to_string(); + input.parse::()?; + match name.as_str() { + "parser" => return Ok(Self::Parser(input.parse::()?)), + _ => panic!("Unrecognized argument {} for option attribute", name), + }; + } + panic!("Arguments to option attribute must be string literals"); + } +} + +pub(crate) fn parse_value_attr(attr: Attribute) -> ValueAttr { + let mut value_attr = ValueAttr::default(); + let Ok(parsed_args) = attr + .parse_args_with(Punctuated::::parse_terminated) + else { + return value_attr; + }; + + for arg in parsed_args { + match arg { + ValueAttrArg::Key(k) => value_attr.keys.push(k), + ValueAttrArg::Value(e) => value_attr.value = Some(e), + }; + } + + value_attr +} + +impl Parse for ValueAttrArg { + fn parse(input: ParseStream) -> syn::Result { + if input.peek(LitStr) { + return Ok(Self::Key(input.parse::()?.value())); + } + + if input.peek(Ident) { + let name = input.parse::()?.to_string(); + input.parse::()?; + match name.as_str() { + "value" => return Ok(Self::Value(input.parse::()?)), + _ => panic!("Unrecognized argument {} for option attribute", name), + }; + } + panic!("Arguments to option attribute must be string literals"); + } +} + +fn parse_flag(input: ParseStream) -> syn::Result { + let str = input.parse::().unwrap().value(); + if let Some(s) = str.strip_prefix("--") { + return Ok(Arg::Long(s.to_owned())); + } else if let Some(s) = str.strip_prefix('-') { + assert_eq!( + s.len(), + 1, + "Exactly one character must follow '-' in a flag attribute" + ); + return Ok(Arg::Short(s.chars().next().unwrap())); + } + panic!("Arguments to flag must start with \"-\" or \"--\""); +} diff --git a/derive/src/lib.rs b/derive/src/lib.rs index 56e0042..0841583 100644 --- a/derive/src/lib.rs +++ b/derive/src/lib.rs @@ -1,40 +1,32 @@ +mod attributes; +use attributes::{ + parse_flag_attr, parse_option_attr, parse_value_attr, FlagAttr, OptionAttr, ValueAttr, +}; + use std::collections::HashMap; use proc_macro::TokenStream; use proc_macro2::TokenStream as TokenStream2; use quote::quote; use syn::{ - parse::{Parse, ParseStream}, - parse_macro_input, - punctuated::Punctuated, - Attribute, - Data::Struct, - DeriveInput, Expr, Fields, Ident, LitStr, Token, + parse_macro_input, Attribute, + Data::{Enum, Struct}, + DeriveInput, Expr, Fields, }; -#[derive(Eq, Hash, PartialEq, Debug)] +#[derive(Eq, Hash, PartialEq, Debug, Clone)] enum Arg { Short(char), Long(String), } -enum OptionsAttribute { - Flag(FlagAttribute), -} - -struct FlagAttribute { - flags: Vec, - value: Option, -} - -enum FlagArg { - Short(char), - Long(String), - Value(Expr), +enum DeriveAttribute { + Flag(FlagAttr), + Option(OptionAttr), } // FIXME: Think of a better name -#[proc_macro_derive(Options, attributes(flag))] +#[proc_macro_derive(Options, attributes(flag, option))] pub fn options(input: TokenStream) -> TokenStream { let input = parse_macro_input!(input as DeriveInput); @@ -58,26 +50,29 @@ pub fn options(input: TokenStream) -> TokenStream { for field in fields.named { let field_ident = field.ident.as_ref().expect("Each field must be named."); let field_name = field_ident.to_string(); - let field_char = field_name.chars().next().unwrap(); for attr in field.attrs { let Some(attr) = parse_attr(attr) else { continue; }; match attr { - OptionsAttribute::Flag(f) => { - let flags = if f.flags.is_empty() { - if field_name.len() > 1 { - vec![Arg::Short(field_char), Arg::Long(field_name.clone())] - } else { - vec![Arg::Short(field_char)] - } - } else { - f.flags - }; - + DeriveAttribute::Flag(f) => { let stmt = match f.value { Some(e) => quote!(self.#field_ident = #e;), None => quote!(self.#field_ident = true;), }; + let flags = flag_names(f.flags, &field_name); + for flag in flags { + map.entry(flag).or_default().push(stmt.clone()); + } + } + DeriveAttribute::Option(o) => { + let stmt = match o.parser { + Some(e) => quote!(self.#field_ident = #e(parser.value()?)?;), + None => { + quote!(self.#field_ident = FromValue::from_value(parser.value()?)?;) + } + }; + + let flags = flag_names(o.flags, &field_name); for flag in flags { map.entry(flag).or_default().push(stmt.clone()); } @@ -102,6 +97,7 @@ pub fn options(input: TokenStream) -> TokenStream { I::Item: Into, { use uutils_args::lexopt; + use uutils_args::FromValue; let mut parser = lexopt::Parser::from_args(args); while let Some(arg) = parser.next()? { match arg { @@ -117,58 +113,86 @@ pub fn options(input: TokenStream) -> TokenStream { TokenStream::from(expanded) } -fn parse_attr(attr: Attribute) -> Option { - if attr.path.is_ident("flag") { - return Some(OptionsAttribute::Flag(parse_flag_attr(attr))); - } - None -} +#[proc_macro_derive(FromValue, attributes(value))] +pub fn from_value(input: TokenStream) -> TokenStream { + let input = parse_macro_input!(input as DeriveInput); -fn parse_flag_attr(attr: Attribute) -> FlagAttribute { - let mut flag_attr = FlagAttribute { - flags: vec![], - value: None, - }; - let Ok(parsed_args) = attr - .parse_args_with(Punctuated::::parse_terminated) - else { - return flag_attr; - }; - for arg in parsed_args { - match arg { - FlagArg::Long(s) => flag_attr.flags.push(Arg::Long(s)), - FlagArg::Short(c) => flag_attr.flags.push(Arg::Short(c)), - FlagArg::Value(e) => flag_attr.value = Some(e), - }; - } - flag_attr -} + // Used in the quasi-quotation below as `#name`. + let name = input.ident; -impl Parse for FlagArg { - fn parse(input: ParseStream) -> syn::Result { - if input.peek(LitStr) { - let str = input.parse::().unwrap().value(); - if let Some(s) = str.strip_prefix("--") { - return Ok(FlagArg::Long(s.to_owned())); - } else if let Some(s) = str.strip_prefix('-') { - assert_eq!( - s.len(), - 1, - "Exactly one character must follow '-' in a flag attribute" - ); - return Ok(FlagArg::Short(s.chars().next().unwrap())); + let (impl_generics, ty_generics, where_clause) = input.generics.split_for_impl(); + + let Enum(data) = input.data else { + panic!("Input should be a struct!"); + }; + + let mut match_arms = vec![]; + for variant in data.variants { + let variant_name = variant.ident.to_string(); + let attrs = variant.attrs.clone(); + for attr in attrs { + if !attr.path.is_ident("value") { + continue; } - panic!("Arguments to flag must start with \"-\" or \"--\""); - } - if input.peek(Ident) { - let name = input.parse::()?.to_string(); - input.parse::()?; - match name.as_str() { - "value" => return Ok(FlagArg::Value(input.parse::()?)), - _ => panic!("Unrecognized argument {} for flag attribute", name), + let ValueAttr { keys, value } = parse_value_attr(attr); + + let keys = if keys.is_empty() { + vec![variant_name.to_lowercase()] + } else { + keys }; + + let stmt = if let Some(v) = value { + quote!(#(| #keys)* => #v) + } else { + let mut v = variant.clone(); + v.attrs = vec![]; + quote!(#(| #keys)* => Self::#v) + }; + match_arms.push(stmt); } - panic!("Arguments to flag attribute must be string literals"); + } + + let expanded = quote!( + impl #impl_generics FromValue for #name #ty_generics #where_clause { + fn from_value(value: std::ffi::OsString) -> Result { + let value = value.into_string()?; + Ok(match value.as_str() { + #(#match_arms),*, + _ => { + return Err(lexopt::Error::ParsingFailed { + value, + error: "Invalid value".into(), + }); + } + }) + } + } + ); + + TokenStream::from(expanded) +} + +fn flag_names(flags: Vec, field_name: &str) -> Vec { + if flags.is_empty() { + let first_char = field_name.chars().next().unwrap(); + if field_name.len() > 1 { + vec![Arg::Short(first_char), Arg::Long(field_name.to_string())] + } else { + vec![Arg::Short(first_char)] + } + } else { + flags + } +} + +fn parse_attr(attr: Attribute) -> Option { + if attr.path.is_ident("flag") { + Some(DeriveAttribute::Flag(parse_flag_attr(attr))) + } else if attr.path.is_ident("option") { + Some(DeriveAttribute::Option(parse_option_attr(attr))) + } else { + None } } diff --git a/src/lib.rs b/src/lib.rs index 74485fd..cf2f60b 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -19,3 +19,19 @@ pub trait Options: Sized + Default { I: IntoIterator + 'static, I::Item: Into; } + +pub trait FromValue: Sized { + fn from_value(value: OsString) -> Result; +} + +impl FromValue for OsString { + fn from_value(value: OsString) -> Result { + Ok(value) + } +} + +impl FromValue for String { + fn from_value(value: OsString) -> Result { + Ok(value.into_string()?) + } +} diff --git a/tests/flags.rs b/tests/flags.rs index 1cadf39..76fbb99 100644 --- a/tests/flags.rs +++ b/tests/flags.rs @@ -1,4 +1,4 @@ -use uutils_args::Options; +use uutils_args::{FromValue, Options}; #[test] fn one_flag() { @@ -255,3 +255,147 @@ fn count() { assert_eq!(Settings::parse(["-vv"]).unwrap().verbosity, 2); assert_eq!(Settings::parse(["-vvv"]).unwrap().verbosity, 3); } + +#[test] +fn string_option() { + #[derive(Default, Options)] + struct Settings { + #[option("--message")] + message: String, + } + + assert_eq!( + Settings::parse(["--message=hello"]).unwrap().message, + "hello" + ); +} + +#[test] +fn enum_option() { + #[derive(FromValue, Default, Debug, PartialEq, Eq)] + enum Format { + #[default] + #[value] + Foo, + #[value] + Bar, + #[value] + Baz, + } + + #[derive(Default, Options)] + struct Settings { + #[option("--format")] + format: Format, + } + + assert_eq!( + Settings::parse(["--format=bar"]).unwrap().format, + Format::Bar + ); + + assert_eq!( + Settings::parse(["--format", "baz"]).unwrap().format, + Format::Baz + ); +} + +#[test] +fn enum_option_with_fields() { + #[derive(FromValue, Default, Debug, PartialEq, Eq)] + enum Indent { + #[default] + Tabs, + #[value("thin", value = Self::Spaces(4))] + #[value("wide", value = Self::Spaces(8))] + Spaces(u8), + } + + #[derive(Default, Options)] + struct Settings { + #[option] + indent: Indent, + } + + assert_eq!( + Settings::parse(["-i=thin"]).unwrap().indent, + Indent::Spaces(4) + ); + assert_eq!( + Settings::parse(["-i=wide"]).unwrap().indent, + Indent::Spaces(8) + ); +} + +#[test] +fn enum_with_complex_from_value() { + #[derive(Default, Debug, PartialEq, Eq)] + enum Indent { + #[default] + Tabs, + Spaces(u8), + } + + impl FromValue for Indent { + fn from_value(value: std::ffi::OsString) -> Result { + let value = value.into_string()?; + if value == "tabs" { + Ok(Self::Tabs) + } else if let Ok(n) = value.parse() { + Ok(Self::Spaces(n)) + } else { + Err(lexopt::Error::ParsingFailed { + value, + error: "Failure!".into(), + }) + } + } + } + + #[derive(Default, Options)] + struct Settings { + #[option] + indent: Indent, + } + + assert_eq!(Settings::parse(["-i=tabs"]).unwrap().indent, Indent::Tabs); + assert_eq!(Settings::parse(["-i=4"]).unwrap().indent, Indent::Spaces(4)); +} + +#[test] +fn color() { + #[derive(Default, FromValue, Debug, PartialEq, Eq)] + enum Color { + #[value("yes", "always")] + Always, + #[default] + #[value("auto")] + Auto, + #[value("no", "never")] + Never, + } + + #[derive(Default, Options)] + struct Settings { + #[option] + color: Color, + } + + assert_eq!( + Settings::parse(["--color=yes"]).unwrap().color, + Color::Always + ); + assert_eq!( + Settings::parse(["--color=always"]).unwrap().color, + Color::Always + ); + assert_eq!(Settings::parse(["--color=no"]).unwrap().color, Color::Never); + assert_eq!( + Settings::parse(["--color=never"]).unwrap().color, + Color::Never + ); + assert_eq!( + Settings::parse(["--color=auto"]).unwrap().color, + Color::Auto + ); +}