use proc_macro2::Span;
use syn::parse::{Parse, ParseStream};
use syn::{Ident, LitInt, Token};
#[derive(Debug, Clone)]
pub enum MathExpr {
Int(i64, Span),
Ident(Ident),
BinOp {
op: BinOp,
lhs: Box<MathExpr>,
rhs: Box<MathExpr>,
},
Neg(Box<MathExpr>),
LogicalNot(Box<MathExpr>),
Func {
name: String,
span: Span,
args: Vec<MathExpr>,
},
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum BinOp {
Add,
Sub,
Mul,
Div,
Pow,
Gt,
Lt,
Ge,
Le,
EqEq,
Ne,
AndAnd,
OrOr,
}
pub const KNOWN_FUNCTIONS: &[&str] = &[
"sin",
"cos",
"tan",
"asin",
"acos",
"atan",
"sinh",
"cosh",
"tanh",
"asinh",
"acosh",
"atanh",
"exp",
"ln",
"sqrt",
"cbrt",
"abs",
"sign",
"floor",
"ceiling",
"sec",
"csc",
"cot",
"acot",
"asec",
"acsc",
"coth",
"sech",
"csch",
"acoth",
"asech",
"acsch",
"sinc",
"arg",
"conjugate",
"fibonacci",
"lucas",
"catalan_number",
"bell",
"euler_number",
"harmonic",
"subfactorial",
"factorial2",
"bernoulli_number",
"heaviside",
"dirac_delta",
"lambertw",
"gamma",
"log_gamma",
"digamma",
"erf",
"erfc",
"beta",
"atan2",
];
pub fn is_known_function(name: &str) -> bool {
KNOWN_FUNCTIONS.contains(&name)
}
pub const KNOWN_CONSTANTS: &[&str] = &["pi", "E", "I", "oo", "nan", "zoo"];
pub fn is_known_constant(name: &str) -> bool {
KNOWN_CONSTANTS.contains(&name)
}
impl MathExpr {
pub fn as_int(&self) -> Option<i64> {
match self {
MathExpr::Int(n, _) => Some(*n),
_ => None,
}
}
pub fn is_numeric_only(&self) -> bool {
match self {
MathExpr::Int(..) => true,
MathExpr::Neg(inner) => inner.is_numeric_only(),
MathExpr::BinOp { op, lhs, rhs } => {
matches!(
op,
BinOp::Add | BinOp::Sub | BinOp::Mul | BinOp::Div | BinOp::Pow
) && lhs.is_numeric_only()
&& rhs.is_numeric_only()
}
_ => false,
}
}
pub fn collect_wilds(&self) -> Vec<Ident> {
let mut wilds = Vec::new();
self.collect_wilds_inner(&mut wilds);
wilds
}
fn collect_wilds_inner(&self, wilds: &mut Vec<Ident>) {
match self {
MathExpr::Ident(id)
if id.to_string().ends_with('_') && !wilds.iter().any(|w| w == id) =>
{
wilds.push(id.clone());
}
MathExpr::BinOp { lhs, rhs, .. } => {
lhs.collect_wilds_inner(wilds);
rhs.collect_wilds_inner(wilds);
}
MathExpr::Neg(inner) => inner.collect_wilds_inner(wilds),
MathExpr::LogicalNot(inner) => inner.collect_wilds_inner(wilds),
MathExpr::Func { args, .. } => {
for arg in args {
arg.collect_wilds_inner(wilds);
}
}
_ => {}
}
}
}
fn infix_bp(op: BinOp) -> (u8, u8) {
match op {
BinOp::OrOr => (1, 2), BinOp::AndAnd => (3, 4), BinOp::EqEq | BinOp::Ne => (5, 6), BinOp::Gt | BinOp::Lt | BinOp::Ge | BinOp::Le => (7, 8), BinOp::Add | BinOp::Sub => (9, 10), BinOp::Mul | BinOp::Div => (11, 12), BinOp::Pow => (16, 15), }
}
fn prefix_bp() -> u8 {
13 }
pub fn parse_math_expr(input: ParseStream) -> syn::Result<MathExpr> {
parse_expr_bp(input, 0)
}
fn parse_expr_bp(input: ParseStream, min_bp: u8) -> syn::Result<MathExpr> {
let mut lhs = parse_prefix(input)?;
while let Some(op) = peek_binop(input) {
let (left_bp, right_bp) = infix_bp(op);
if left_bp < min_bp {
break; }
consume_binop(input, op)?;
let rhs = parse_expr_bp(input, right_bp)?;
lhs = MathExpr::BinOp {
op,
lhs: Box::new(lhs),
rhs: Box::new(rhs),
};
}
Ok(lhs)
}
fn parse_prefix(input: ParseStream) -> syn::Result<MathExpr> {
if input.peek(Token![-]) {
let _: Token![-] = input.parse()?;
let operand = parse_expr_bp(input, prefix_bp())?;
Ok(MathExpr::Neg(Box::new(operand)))
} else if input.peek(Token![!]) {
let _: Token![!] = input.parse()?;
let operand = parse_expr_bp(input, prefix_bp())?;
Ok(MathExpr::LogicalNot(Box::new(operand)))
} else {
parse_primary(input)
}
}
fn parse_primary(input: ParseStream) -> syn::Result<MathExpr> {
if input.peek(syn::token::Paren) {
let content;
syn::parenthesized!(content in input);
parse_math_expr(&content)
} else if input.peek(LitInt) {
let lit: LitInt = input.parse()?;
let value: i64 = lit.base10_parse()?;
Ok(MathExpr::Int(value, lit.span()))
} else if input.peek(Ident) {
let ident: Ident = input.parse()?;
let name = ident.to_string();
if input.peek(syn::token::Paren) {
let content;
syn::parenthesized!(content in input);
let args = parse_arg_list(&content)?;
Ok(MathExpr::Func {
name,
span: ident.span(),
args,
})
} else {
Ok(MathExpr::Ident(ident))
}
} else {
Err(input.error("expected integer, identifier, function call, or '('"))
}
}
fn parse_arg_list(input: ParseStream) -> syn::Result<Vec<MathExpr>> {
let mut args = Vec::new();
if input.is_empty() {
return Ok(args);
}
args.push(parse_math_expr(input)?);
while input.peek(Token![,]) {
let _: Token![,] = input.parse()?;
args.push(parse_math_expr(input)?);
}
Ok(args)
}
fn peek_binop(input: ParseStream) -> Option<BinOp> {
if input.peek(Token![&&]) {
return Some(BinOp::AndAnd);
}
if input.peek(Token![||]) {
return Some(BinOp::OrOr);
}
if input.peek(Token![>=]) {
return Some(BinOp::Ge);
}
if input.peek(Token![<=]) {
return Some(BinOp::Le);
}
if input.peek(Token![==]) {
return Some(BinOp::EqEq);
}
if input.peek(Token![!=]) {
return Some(BinOp::Ne);
}
if input.peek(Token![>]) {
return Some(BinOp::Gt);
}
if input.peek(Token![<]) {
return Some(BinOp::Lt);
}
if input.peek(Token![+]) {
Some(BinOp::Add)
} else if input.peek(Token![-]) {
Some(BinOp::Sub)
} else if input.peek(Token![*]) {
Some(BinOp::Mul)
} else if input.peek(Token![/]) {
Some(BinOp::Div)
} else if input.peek(Token![^]) {
Some(BinOp::Pow)
} else {
None
}
}
fn consume_binop(input: ParseStream, op: BinOp) -> syn::Result<()> {
match op {
BinOp::Add => {
let _: Token![+] = input.parse()?;
}
BinOp::Sub => {
let _: Token![-] = input.parse()?;
}
BinOp::Mul => {
let _: Token![*] = input.parse()?;
}
BinOp::Div => {
let _: Token![/] = input.parse()?;
}
BinOp::Pow => {
let _: Token![^] = input.parse()?;
}
BinOp::Gt => {
let _: Token![>] = input.parse()?;
}
BinOp::Lt => {
let _: Token![<] = input.parse()?;
}
BinOp::Ge => {
let _: Token![>=] = input.parse()?;
}
BinOp::Le => {
let _: Token![<=] = input.parse()?;
}
BinOp::EqEq => {
let _: Token![==] = input.parse()?;
}
BinOp::Ne => {
let _: Token![!=] = input.parse()?;
}
BinOp::AndAnd => {
let _: Token![&&] = input.parse()?;
}
BinOp::OrOr => {
let _: Token![||] = input.parse()?;
}
}
Ok(())
}
pub struct ExprMacroInput {
pub ctx: Ident,
pub expr: MathExpr,
}
impl Parse for ExprMacroInput {
fn parse(input: ParseStream) -> syn::Result<Self> {
let ctx: Ident = input.parse()?;
input.parse::<Token![,]>()?;
let expr = parse_math_expr(input)?;
Ok(ExprMacroInput { ctx, expr })
}
}
pub struct RuleMacroInput {
pub arena: Ident,
pub name: syn::LitStr,
pub lhs: MathExpr,
pub rhs: MathExpr,
pub condition: Option<syn::Expr>,
}
impl Parse for RuleMacroInput {
fn parse(input: ParseStream) -> syn::Result<Self> {
let arena: Ident = input.parse()?;
let _: Token![,] = input.parse()?;
let name: syn::LitStr = input.parse()?;
let _: Token![,] = input.parse()?;
let lhs = parse_math_expr(input)?;
let _: Token![=>] = input.parse()?;
let rhs = parse_math_expr(input)?;
let condition = if input.peek(Token![if]) {
input.parse::<Token![if]>()?;
Some(input.parse::<syn::Expr>()?)
} else {
None
};
Ok(RuleMacroInput {
arena,
name,
lhs,
rhs,
condition,
})
}
}
pub struct MatrixMacroInput {
pub ctx: Ident,
pub rows: Vec<Vec<MathExpr>>,
}
impl Parse for MatrixMacroInput {
fn parse(input: ParseStream) -> syn::Result<Self> {
let ctx: Ident = input.parse()?;
input.parse::<Token![,]>()?;
let mut rows = Vec::new();
while !input.is_empty() {
let row_content;
syn::bracketed!(row_content in input);
let mut row = Vec::new();
loop {
row.push(parse_math_expr(&row_content)?);
if row_content.is_empty() {
break;
}
row_content.parse::<Token![,]>()?;
if row_content.is_empty() {
break; }
}
rows.push(row);
if input.is_empty() {
break;
}
if input.peek(Token![,]) {
input.parse::<Token![,]>()?;
}
}
if rows.is_empty() {
return Err(syn::Error::new(
Span::call_site(),
"matrix! needs at least one row: `matrix![ctx, [a, b], [c, d]]`",
));
}
if let Some(first_len) = rows.first().map(|r| r.len()) {
for (i, row) in rows.iter().enumerate() {
if row.len() != first_len {
return Err(syn::Error::new(
Span::call_site(),
format!(
"matrix! row {} has {} columns, but row 0 has {} columns",
i,
row.len(),
first_len
),
));
}
}
}
Ok(MatrixMacroInput { ctx, rows })
}
}
pub struct EqMacroInput {
pub ctx: Ident,
pub lhs: MathExpr,
pub rhs: MathExpr,
}
impl Parse for EqMacroInput {
fn parse(input: ParseStream) -> syn::Result<Self> {
let ctx: Ident = input.parse()?;
input.parse::<Token![,]>()?;
let lhs = parse_math_expr(input)?;
input.parse::<Token![=]>()?;
let rhs = parse_math_expr(input)?;
Ok(EqMacroInput { ctx, lhs, rhs })
}
}
pub struct DimMacroInput {
pub ctx: Ident,
pub output_type: syn::Type,
pub expr: MathExpr,
}
impl Parse for DimMacroInput {
fn parse(input: ParseStream) -> syn::Result<Self> {
let ctx: Ident = input.parse()?;
input.parse::<Token![,]>()?;
let output_type: syn::Type = input.parse()?;
input.parse::<Token![:]>()?;
let expr = parse_math_expr(input)?;
Ok(DimMacroInput {
ctx,
output_type,
expr,
})
}
}