#![cfg_attr(debug_assertions, allow(unused_imports))]
use proc_macro::{self, TokenStream};
use proc_macro_error::abort;
use proc_macro2::{Ident, Span, TokenStream as TokenStream2};
use quote::quote;
use syn::ext::IdentExt;
use syn::parse::Parse;
use syn::{Token, parse_macro_input};
#[proc_macro]
pub fn compare_variables(input: TokenStream) -> TokenStream {
let comparison_error_info: ComparisonErrorInfo = parse_macro_input!(input);
let first_arg = comparison_error_info.first_arg.as_token_stream();
let relation_first_to_second = comparison_error_info
.relation_first_to_second
.as_token_stream();
let second_arg = comparison_error_info.second_arg.as_token_stream();
let relation_second_to_third = comparison_error_info
.relation_second_to_third
.as_token_stream();
let third_arg = match comparison_error_info.third_arg {
Some(arg) => {
let ts = arg.as_token_stream();
quote! {Some(#ts)}
}
None => quote! {None},
};
let stream = quote! {
compare_variables::Comparison::new_checked(
#first_arg,
#relation_first_to_second,
#second_arg,
#relation_second_to_third,
#third_arg,
)
};
return TokenStream::from(stream);
}
#[repr(u8)]
enum Operator {
Lesser,
LesserOrEqual,
Equal,
Inequal,
GreaterOrEqual,
Greater,
}
impl Operator {
fn as_token_stream(&self) -> proc_macro2::TokenStream {
match self {
Operator::Lesser => {
quote! {
compare_variables::ComparisonOperator::Lesser
}
}
Operator::LesserOrEqual => {
quote! {
compare_variables::ComparisonOperator::LesserOrEqual
}
}
Operator::Equal => {
quote! {
compare_variables::ComparisonOperator::Equal
}
}
Operator::Inequal => {
quote! {
compare_variables::ComparisonOperator::Inequal
}
}
Operator::GreaterOrEqual => {
quote! {
compare_variables::ComparisonOperator::GreaterOrEqual
}
}
Operator::Greater => {
quote! {
compare_variables::ComparisonOperator::Greater
}
}
}
}
}
impl Parse for Operator {
fn parse(input: syn::parse::ParseStream) -> syn::Result<Self> {
if input.peek(Token![<=]) {
input.parse::<Token![<=]>()?;
Ok(Operator::LesserOrEqual)
} else if input.peek(Token![>=]) {
input.parse::<Token![>=]>()?;
Ok(Operator::GreaterOrEqual)
} else if input.peek(Token![==]) {
input.parse::<Token![==]>()?;
Ok(Operator::Equal)
} else if input.peek(Token![!=]) {
input.parse::<Token![!=]>()?;
Ok(Operator::Inequal)
} else if input.peek(Token![<]) {
input.parse::<Token![<]>()?;
Ok(Operator::Lesser)
} else if input.peek(Token![>]) {
input.parse::<Token![>]>()?;
Ok(Operator::Greater)
} else {
Err(syn::Error::new(
input.span(),
"no comparison operator could be identified. Valid
operators are \"<\", \"<=\", \"==\", \"!=\", \">=\" or \">\".",
))
}
}
}
enum VariableOrLiteral {
Other {
arg_names: Vec<String>,
arg_names_display: Vec<String>,
},
LitFloat(syn::LitFloat),
LitInt(syn::LitInt),
}
impl VariableOrLiteral {
fn as_token_stream(&self) -> proc_macro2::TokenStream {
match self {
VariableOrLiteral::Other {
arg_names,
arg_names_display,
} => {
let arg_value = arg_names.join(".");
let arg_value_ts: TokenStream2 = match str::parse::<TokenStream2>(&arg_value) {
Ok(ts) => ts,
Err(_) => abort!(
Span::call_site(),
format!("could not interpret {arg_value} as rust code")
),
};
if arg_names_display.is_empty() {
quote! {
compare_variables::ComparisonValue::new(#arg_value_ts, None)
}
} else {
let arg_name_display = arg_names_display.join(".");
quote! {
compare_variables::ComparisonValue::new(#arg_value_ts, Some(#arg_name_display))
}
}
}
VariableOrLiteral::LitFloat(lit) => {
quote! {
compare_variables::ComparisonValue::new(#lit, None)
}
}
VariableOrLiteral::LitInt(lit) => {
quote! {
compare_variables::ComparisonValue::new(#lit, None)
}
}
}
}
}
impl Parse for VariableOrLiteral {
fn parse(input: syn::parse::ParseStream) -> syn::Result<Self> {
fn parse_composite_varname(
input: &syn::parse::ParseStream,
vec: &mut Vec<String>,
) -> syn::Result<()> {
loop {
if input.peek(syn::LitInt) {
let lit = input.parse::<syn::LitInt>()?;
vec.push(lit.to_string());
} else {
let ident: syn::Ident = input.call(Ident::parse_any)?;
vec.push(ident.to_string()); }
if input.peek(Token![.]) {
let _ = input.parse::<Token![.]>()?;
} else {
break;
}
}
return Ok(());
}
if input.peek(syn::LitFloat) {
let val = input.parse::<syn::LitFloat>()?;
return Ok(VariableOrLiteral::LitFloat(val));
} else if input.peek(syn::LitInt) {
let val = input.parse::<syn::LitInt>()?;
return Ok(VariableOrLiteral::LitInt(val));
} else {
let mut display_arg_names = true;
let mut arg_names: Vec<String> = Vec::new();
let first_ident: Ident = input.call(Ident::parse_any)?;
if input.peek(Token![.]) {
let _ = input.parse::<Token![.]>()?;
arg_names.push(first_ident.to_string());
parse_composite_varname(&input, &mut arg_names)?;
} else {
if input.peek(syn::Ident) {
if first_ident == "val" {
display_arg_names = false;
parse_composite_varname(&input, &mut arg_names)?;
} else {
abort!(
Span::call_site(),
format!("found unexpected tokens behind {first_ident}")
)
}
} else {
arg_names.push(first_ident.to_string());
}
}
let arg_names_display: Vec<String> = if input.peek(Token![as]) {
input.parse::<Token![as]>()?;
let mut arg_names_display: Vec<String> = Vec::new();
parse_composite_varname(&input, &mut arg_names_display)?;
if display_arg_names {
arg_names_display
} else {
Vec::new()
}
} else {
if display_arg_names {
arg_names.clone()
} else {
Vec::new()
}
};
return Ok(VariableOrLiteral::Other {
arg_names,
arg_names_display,
});
}
}
}
struct ComparisonErrorInfo {
first_arg: VariableOrLiteral,
relation_first_to_second: Operator,
second_arg: VariableOrLiteral,
relation_second_to_third: Operator,
third_arg: Option<VariableOrLiteral>,
}
impl Parse for ComparisonErrorInfo {
fn parse(input: syn::parse::ParseStream) -> syn::Result<Self> {
let first_arg = VariableOrLiteral::parse(&input)?;
let relation_first_to_second = Operator::parse(&input)?;
let second_arg = VariableOrLiteral::parse(&input)?;
let (relation_second_to_third, third_arg) = if let Ok(operator) = Operator::parse(&input) {
(operator, Some(VariableOrLiteral::parse(&input)?))
} else {
(Operator::Equal, None)
};
return Ok(ComparisonErrorInfo {
first_arg,
relation_first_to_second,
second_arg,
relation_second_to_third,
third_arg,
});
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_parse_check_bounds_info() {
let _: ComparisonErrorInfo = syn::parse_quote!(0.0 < arg);
let _: ComparisonErrorInfo = syn::parse_quote!(0.0 <= arg);
let _: ComparisonErrorInfo = syn::parse_quote!(0.0 <= arg as alternative_arg);
let _: ComparisonErrorInfo = syn::parse_quote!(0.0 < arg <= 1.0);
let _: ComparisonErrorInfo = syn::parse_quote!(0.0 < arg as alternative_arg <= 1.0);
let _: ComparisonErrorInfo = syn::parse_quote!(arg < 1.0);
let _: ComparisonErrorInfo = syn::parse_quote!(arg <= 1.0);
let _: ComparisonErrorInfo = syn::parse_quote!(arg as alternative_arg <= 1.0);
let _: ComparisonErrorInfo = syn::parse_quote!(-1 < arg);
let _: ComparisonErrorInfo = syn::parse_quote!(-1 < -2);
let _: ComparisonErrorInfo = syn::parse_quote!(-1 < arg as alternative_arg <= 2);
}
}