You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

Rust中无需字符串化解析访问令牌树/表达式令牌的方法及CAS实现咨询

实现Rust宏驱动的符号代数系统(CAS)

要实现你需求中的宏驱动CAS,核心是利用Rust的过程宏直接操作语法令牌树,结合syn/quote库解析和生成代码,无需字符串化重解析。以下是具体实现步骤:

1. 依赖准备

在Cargo.toml中添加过程宏相关依赖:

[dependencies]
syn = { version = "2.0", features = ["full"] }
quote = "1.0"
proc-macro2 = "1.0"

[lib]
proc-macro = true

2. 定义宏的输入解析结构

用syn解析symbolic!宏中的函数定义,提取函数名、参数和表达式:

use syn::{parse_macro_input, ItemFn, Ident, PatType, Expr, Parse, ParseStream};

struct SymbolicFunctionDef {
    name: Ident,
    params: Vec<Ident>,
    body: Expr,
}

impl Parse for SymbolicFunctionDef {
    fn parse(input: ParseStream) -> syn::Result<Self> {
        let fn_item: ItemFn = input.parse()?;
        let params = fn_item.sig.inputs.iter()
            .filter_map(|arg| match arg {
                syn::FnArg::Typed(PatType { pat, .. }) => match **pat {
                    syn::Pat::Ident(ref ident_pat) => Some(ident_pat.ident.clone()),
                    _ => None,
                },
                _ => None,
            })
            .collect();
        let body = match fn_item.block.stmts.first() {
            Some(syn::Stmt::Expr(expr)) => expr.clone(),
            Some(syn::Stmt::Semi(expr, _)) => expr.clone(),
            _ => return Err(syn::Error::new(fn_item.block.span, "函数体必须是单个表达式")),
        };
        Ok(SymbolicFunctionDef {
            name: fn_item.sig.ident,
            params,
            body,
        })
    }
}

3. 实现表达式到AST的转换

通过syn的访问者模式,遍历表达式令牌,将其转换为你定义的Symbolic和Node类型:

use syn::visit_mut::VisitMut;
use proc_macro2::TokenStream;
use quote::quote;

struct AstConverter;

impl VisitMut for AstConverter {
    fn visit_ident_mut(&mut self, i: &mut Ident) {
        // 将标识符转换为Symbolic::Symbol
        let ident_str = i.to_string();
        *i = syn::parse2(quote!(Symbolic::Symbol(#ident_str))).unwrap();
    }

    fn visit_lit_int_mut(&mut self, lit: &mut syn::LitInt) {
        // 将整数字面量转换为Symbolic::Rational
        let num = lit.base10_parse::<isize>().unwrap();
        *lit = syn::parse2(quote!(Symbolic::Rational(#num, 1))).unwrap();
    }

    fn visit_expr_binary_mut(&mut self, expr: &mut syn::ExprBinary) {
        // 递归处理左右操作数
        self.visit_expr_mut(&mut expr.left);
        self.visit_expr_mut(&mut expr.right);

        // 替换运算符为对应的Node构造
        let op_token = match &expr.op {
            syn::BinOp::Add(_) => quote!('+'),
            syn::BinOp::Sub(_) => quote!('-'),
            syn::BinOp::Mul(_) => quote!('*'),
            syn::BinOp::Div(_) => quote!('/'),
            syn::BinOp::Caret(_) => quote!('^'),
            _ => panic!("不支持的运算符"),
        };

        let lhs = &expr.left;
        let rhs = &expr.right;
        *expr = syn::parse2(quote!(Node::BinaryExpr { op: #op_token, lhs: Box::new(#lhs), rhs: Box::new(#rhs) })).unwrap();
    }
}

4. 宏展开逻辑

编写过程宏,将输入的函数定义转换为包含AST的SymbolicFunction实例:

#[proc_macro]
pub fn symbolic(input: TokenStream) -> TokenStream {
    let def = parse_macro_input!(input as SymbolicFunctionDef);
    let mut body = def.body;
    let mut converter = AstConverter;
    converter.visit_expr_mut(&mut body);

    let name = def.name;
    let params = def.params.iter().map(|p| p.to_string());
    quote! {
        let #name = SymbolicFunction {
            params: vec![#(#params),*],
            ast: #body,
        };
    }.into()
}

5. 完善运算符重载与求导逻辑

补充Symbolic的运算符重载,让表达式能自动构建Node:

use std::ops::{Add, Mul, Div, Sub, Pow};

impl Mul<Self> for Symbolic {
    type Output = Node;
    fn mul(self, rhs: Self) -> Self::Output {
        Node::BinaryExpr {
            op: '*',
            lhs: Box::new(Node::Symb(self)),
            rhs: Box::new(Node::Symb(rhs)),
        }
    }
}

impl Div<Self> for Symbolic {
    type Output = Node;
    fn div(self, rhs: Self) -> Self::Output {
        Node::BinaryExpr {
            op: '/',
            lhs: Box::new(Node::Symb(self)),
            rhs: Box::new(Node::Symb(rhs)),
        }
    }
}

impl Pow<Self> for Symbolic {
    type Output = Node;
    fn pow(self, rhs: Self) -> Self::Output {
        Node::BinaryExpr {
            op: '^',
            lhs: Box::new(Node::Symb(self)),
            rhs: Box::new(Node::Symb(rhs)),
        }
    }
}

// 补充Sub等其他运算符重载...

然后给Node实现递归求导方法:

impl Node {
    pub fn derive(&self, var: &str) -> Node {
        match self {
            Node::Symb(sym) => match sym {
                Symbolic::Symbol(name) if name == var => Node::Symb(Symbolic::Rational(1, 1)),
                Symbolic::Rational(_, _) => Node::Symb(Symbolic::Rational(0, 1)),
                _ => Node::Symb(Symbolic::Rational(0, 1)),
            },
            Node::BinaryExpr { op, lhs, rhs } => match op {
                '+' => Node::BinaryExpr {
                    op: '+',
                    lhs: Box::new(lhs.derive(var)),
                    rhs: Box::new(rhs.derive(var)),
                },
                '*' => Node::BinaryExpr {
                    op: '+',
                    lhs: Box::new(Node::BinaryExpr {
                        op: '*',
                        lhs: Box::new(lhs.derive(var)),
                        rhs: rhs.clone(),
                    }),
                    rhs: Box::new(Node::BinaryExpr {
                        op: '*',
                        lhs: lhs.clone(),
                        rhs: Box::new(rhs.derive(var)),
                    }),
                },
                '^' => {
                    // 处理常数幂的求导,变量幂需额外逻辑
                    let (num, den) = match &**rhs {
                        Node::Symb(Symbolic::Rational(n, d)) => (*n, *d),
                        _ => panic!("暂不支持变量幂求导"),
                    };
                    Node::BinaryExpr {
                        op: '*',
                        lhs: Box::new(Node::Symb(Symbolic::Rational(num, den))),
                        rhs: Box::new(Node::BinaryExpr {
                            op: '*',
                            lhs: Box::new(Node::BinaryExpr {
                                op: '^',
                                lhs: lhs.clone(),
                                rhs: Box::new(Node::Symb(Symbolic::Rational(num - 1, den))),
                            }),
                            rhs: Box::new(lhs.derive(var)),
                        }),
                    }
                }
                '/' => todo!("实现除法求导法则"),
                '-' => todo!("实现减法求导法则"),
                _ => panic!("不支持的运算符求导"),
            },
            Node::UnaryExpr { op, child } => todo!("实现一元表达式求导"),
        }
    }
}

// 给SymbolicFunction封装derive方法
struct SymbolicFunction {
    params: Vec<String>,
    ast: Node,
}

impl SymbolicFunction {
    pub fn derive(&self, var: &str) -> Node {
        self.ast.derive(var)
    }
}

6. 测试使用

现在可以按你预期的方式使用宏:

symbolic!(fn function(x, a) -> 2/4*x^2 + a*x + 4);
let derivative = function.derive("x");
// derivative应为x + a(简化2/4为1/2后,1/2*2x = x,a*x求导为a)

内容的提问来源于stack exchange,提问作者Lyndon Alcock

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.10 10:50:23