如何为含自身指针的AST节点类编写相等运算符?
Great question—let's break this down step by step, since your operator== is key to making common subexpression elimination (CSE) work correctly.
首先:注释里的递归调用完全可行(而且是必须的)
你注释掉的那些判断,完全可以递归调用operator==,这正是遍历并比较整个AST树结构的正确方式。因为你的子节点left和right本身就是ExpNode实例,递归调用能自动深入到树的每一层,确保两个表达式树的所有节点都匹配,而不是只比较顶层的子节点。
比如,对于两个加法操作的节点,它们的子节点可能也是操作节点,递归调用*left == *rhs.left就能自动比较那些子树的结构,这比手动一层层写判断要简洁且不易出错。
但当前代码存在几个关键问题,会导致判断错误
你的现有实现有几个漏洞,会让CSE逻辑失效甚至崩溃:
1. 未处理空指针风险
如果某个节点的left或right是nullptr(比如单目操作,或者节点初始化不完整),直接访问left->token会触发崩溃。必须先检查指针是否为空。
2. 未检查当前节点的核心属性
比如当前节点是operation类型时,你完全没比较op或opname!这会把a + b和a * b当成相等的表达式,这显然是错的。
3. 交换子节点的逻辑不能适用于所有操作
你的代码默认所有操作都满足交换律(允许左右子节点交换后仍判定相等),但像减法a - b、除法a / b、幂运算a^b这些操作是不满足交换律的,交换后是完全不同的表达式,不能判定为相等。
4. 缺少默认返回值
如果所有分支都不匹配,你的函数没有返回任何值,这会导致未定义行为(编译器可能返回随机值)。
推荐的改进实现
下面是一个更健壮的operator==版本,解决了上述问题,且正确递归比较子树:
bool operator==(const ExpNode& rhs) const{ // 先比较当前节点的token,token不同直接不相等 if (token != rhs.token) { return false; } // 根据token类型分别处理 switch(token) { case constant: return std::abs(value - rhs.value) < 1e-9; // 浮点数比较用误差范围 case input: return name == rhs.name; case operation: // 先比较操作符,操作符不同直接不相等 if (op != rhs.op || opname != rhs.opname) { return false; } // 检查子节点是否为空的情况 bool leftEqual = false; bool rightEqual = false; // 处理左子节点:都为空 或者 都不为空且相等 if (left == nullptr && rhs.left == nullptr) { leftEqual = true; } else if (left != nullptr && rhs.left != nullptr) { leftEqual = *left == *rhs.left; } // 处理右子节点:都为空 或者 都不为空且相等 if (right == nullptr && rhs.right == nullptr) { rightEqual = true; } else if (right != nullptr && rhs.right != nullptr) { rightEqual = *right == *rhs.right; } // 先判断正常顺序是否相等 if (leftEqual && rightEqual) { return true; } // 如果是满足交换律的操作,再判断交换子节点的情况 bool isCommutative = (op == '+' || op == '*'); // 根据你的操作符定义调整 if (isCommutative) { bool swappedLeftEqual = false; bool swappedRightEqual = false; if (left != nullptr && rhs.right != nullptr) { swappedLeftEqual = *left == *rhs.right; } else if (left == nullptr && rhs.right == nullptr) { swappedLeftEqual = true; } if (right != nullptr && rhs.left != nullptr) { swappedRightEqual = *right == *rhs.left; } else if (right == nullptr && rhs.left == nullptr) { swappedRightEqual = true; } return swappedLeftEqual && swappedRightEqual; } // 非交换律操作,正常顺序不相等就返回false return false; default: return false; } }
额外的优化建议
- 浮点数比较:用
std::abs(value - rhs.value) < 1e-9代替直接相等,避免精度误差导致的误判。 - 空指针检查:必须在访问子节点的任何成员之前,判断指针是否为空,否则会触发未定义行为。
- 交换律操作的扩展:如果你的支持更多交换律操作(比如按位与、按位或),可以扩展
isCommutative的判断条件。
内容的提问来源于stack exchange,提问作者Endre

