diff --git a/src/sindi/__init__.py b/src/sindi/__init__.py index 2e7d975..401bfc6 100644 --- a/src/sindi/__init__.py +++ b/src/sindi/__init__.py @@ -5,6 +5,7 @@ from .tokenizer import Tokenizer from .parser import Parser, ASTNode from .simplifier import Simplifier +from .ast_rewriter import ASTRewriter __all__ = [ "Comparator", @@ -14,6 +15,7 @@ "Parser", "ASTNode", "Simplifier", + "ASTRewriter", ] __version__ = "0.2.0" diff --git a/src/sindi/ast_rewriter.py b/src/sindi/ast_rewriter.py new file mode 100644 index 0000000..1399894 --- /dev/null +++ b/src/sindi/ast_rewriter.py @@ -0,0 +1,269 @@ +# src/sindi/ast_rewriter.py +from __future__ import annotations +from typing import List, Optional, Tuple +from .parser import ASTNode + +COMMUTATIVE_BOOL = {"&&", "||", "==", "!="} +ASSOCIATIVE = {"&&", "||", "+"} +COMMUTATIVE_ARITH = {"+", "*"} + +REL_OPS = {">", ">=", "<", "<=", "==", "!="} + + +def _clone(n: ASTNode) -> ASTNode: + return ASTNode(n.value, [_clone(c) for c in n.children]) + + +def _repr(n: ASTNode) -> str: + return repr(n) + + +def _eq(a: ASTNode, b: ASTNode) -> bool: + return _repr(a) == _repr(b) + + +def _sort_children(node: ASTNode) -> None: + node.children.sort(key=_repr) + + +def _is_bool_leaf(n: ASTNode) -> bool: + return (not n.children) and (n.value.lower() in ("true", "false")) + + +def _is_zero(n: ASTNode) -> bool: + if n.children: + return False + try: + return int(n.value) == 0 + except Exception: + try: + return float(n.value) == 0.0 + except Exception: + return False + + +def _is_number(n: ASTNode) -> bool: + if n.children: + return False + try: + int(n.value) + return True + except Exception: + try: + float(n.value) + return True + except Exception: + return False + + +class ASTRewriter: + """ + Post-parse AST canonicalizer. + + This pass works AFTER string-level rewriting + tokenization + parsing, + so it has full structure. It focuses on: + - Boolean-equality folding: (X == true) → X, (X == false) → !X, etc. + - NOT cleanups: !!X → X + - Simple algebraic normalization across inequalities: + X < (Y - Z) -> (X + Z) < Y + (A - B) < C -> A < (C + B) + - Flattening associative ops (&&, ||, +) + - Sorting commutative children (&&, ||, +, ==, !=) for determinism + - Reordering equality arguments deterministically + - Plus-commutativity/coalescing + - A few semantic patterns (owner/admin convenience, bitmask-finalized) + """ + + # ---- convenience substitutions you previously had as strings ---- + # Keep these here as AST-level forms in case the surface pass didn't catch them. + def _owner_admin_forms(self, n: ASTNode) -> ASTNode: + # now -> block.timestamp (parsers emit 'now' as a plain IDENTIFIER leaf) + if not n.children and n.value == "now": + return ASTNode("block.timestamp") + + # isOwner() -> (msg.sender == owner()) + if n.value == "isOwner()" and not n.children: + return ASTNode("==", [ASTNode("msg.sender"), ASTNode("owner()")]) + + # isAdmin() -> (msg.sender == admin) + if n.value == "isAdmin()" and not n.children: + return ASTNode("==", [ASTNode("msg.sender"), ASTNode("admin")]) + + # _msgSender() -> msg.sender + if n.value == "_msgSender()" and not n.children: + return ASTNode("msg.sender") + + return n + + def _flatten(self, node: ASTNode) -> ASTNode: + if not node.children: + return node + node.children = [self._flatten(c) for c in node.children] + if node.value in ASSOCIATIVE: + flat: List[ASTNode] = [] + for ch in node.children: + if ch.value == node.value: + flat.extend(ch.children) + else: + flat.append(ch) + node.children = flat + return node + + def _comm_sort(self, node: ASTNode) -> ASTNode: + if not node.children: + return node + node.children = [self._comm_sort(c) for c in node.children] + if node.value in (COMMUTATIVE_BOOL | COMMUTATIVE_ARITH): + _sort_children(node) + return node + + def _normalize_equals_to_bool(self, node: ASTNode) -> ASTNode: + """ + X == true -> X + X == false -> !X + X != true -> !X + X != false -> X + (and symmetric forms where true/false is on left) + """ + if not node.children: + return node + node.children = [self._normalize_equals_to_bool(c) for c in node.children] + + if node.value in ("==", "!=") and len(node.children) == 2: + L, R = node.children + + # normalize: (expr op bool) + if _is_bool_leaf(L) or _is_bool_leaf(R): + expr = R if _is_bool_leaf(L) else L + blf = L if _is_bool_leaf(L) else R + bval = blf.value.lower() == "true" + + if node.value == "==": + return expr if bval else ASTNode("!", [expr]) + else: # "!=" + return ASTNode("!", [expr]) if bval else expr + return node + + def _normalize_nots(self, node: ASTNode) -> ASTNode: + if not node.children: + return node + node.children = [self._normalize_nots(c) for c in node.children] + if node.value == "!" and len(node.children) == 1: + ch = node.children[0] + if ch.value == "!": + return ch.children[0] + return node + + def _normalize_rel_sub(self, node: ASTNode) -> ASTNode: + """ + Move a single '-' across inequalities/equalities: + X < Y - Z → X + Z < Y + A - B < C → A < C + B + Similar for <=, >, >=, ==, != + """ + if not node.children: + return node + node.children = [self._normalize_rel_sub(c) for c in node.children] + + if node.value in REL_OPS and len(node.children) == 2: + L, R = node.children + + # Right = (A - B) → (L + B) op A + if R.value == "-" and len(R.children) == 2: + A, B = R.children + return ASTNode(node.value, [ASTNode("+", [L, B]), A]) + + # Left = (A - B) → A op (R + B) + if L.value == "-" and len(L.children) == 2: + A, B = L.children + return ASTNode(node.value, [A, ASTNode("+", [R, B])]) + + return node + + def _normalize_plus_comm(self, node: ASTNode) -> ASTNode: + if not node.children: + return node + node.children = [self._normalize_plus_comm(c) for c in node.children] + node = self._flatten(node) + if node.value == "+": + _sort_children(node) + return node + + def _normalize_commutative_rel_args(self, node: ASTNode) -> ASTNode: + """ + Reorder args for == and != deterministically: (smaller repr) first. + """ + if not node.children: + return node + node.children = [self._normalize_commutative_rel_args(c) for c in node.children] + if node.value in ("==", "!=") and len(node.children) == 2: + L, R = node.children + if _repr(R) < _repr(L): + node.children = [R, L] + return node + + def _finalized_bitmask(self, node: ASTNode) -> ASTNode: + """ + (X & MarketplaceLib.FLAG_MASK_FINALIZED) == 0 → !MarketplaceLib.isFinalized(X) + Handle symmetry too: 0 == (X & FLAG) + """ + if not node.children: + return node + node.children = [self._finalized_bitmask(c) for c in node.children] + + def _mk_is_finalized(x: ASTNode) -> ASTNode: + return ASTNode("!", [ASTNode("MarketplaceLib.isFinalized()", [x])]) + + if node.value == "==" and len(node.children) == 2: + L, R = node.children + + # Left (& ...), Right 0 + if L.value == "&" and _is_zero(R) and len(L.children) == 2: + a, b = L.children + if (not b.children) and b.value == "MarketplaceLib.FLAG_MASK_FINALIZED": + return _mk_is_finalized(a) + if (not a.children) and a.value == "MarketplaceLib.FLAG_MASK_FINALIZED": + return _mk_is_finalized(b) + + # Symmetric: Left 0, Right (& ...) + if _is_zero(L) and R.value == "&" and len(R.children) == 2: + a, b = R.children + if (not b.children) and b.value == "MarketplaceLib.FLAG_MASK_FINALIZED": + return _mk_is_finalized(a) + if (not a.children) and a.value == "MarketplaceLib.FLAG_MASK_FINALIZED": + return _mk_is_finalized(b) + + return node + + def _apply_local_node_rules(self, node: ASTNode) -> ASTNode: + """ + Node-local translations that don't need a global traversal context. + """ + node = self._owner_admin_forms(node) + return node + + # -------- Pipeline -------- + def normalize(self, root: ASTNode) -> ASTNode: + # Work on a clone to keep caller's tree untouched + n = _clone(root) + + # Local substitutions on each node (single-visit) + def _walk_apply(n: ASTNode) -> ASTNode: + if not n.children: + return self._apply_local_node_rules(n) + n.children = [_walk_apply(c) for c in n.children] + return self._apply_local_node_rules(n) + + n = _walk_apply(n) + + # Structured normalizations (multi-pass safe order) + n = self._normalize_equals_to_bool(n) + n = self._normalize_nots(n) + n = self._normalize_rel_sub(n) + n = self._flatten(n) + n = self._normalize_plus_comm(n) + n = self._comm_sort(n) + n = self._normalize_commutative_rel_args(n) + n = self._finalized_bitmask(n) + + return n \ No newline at end of file diff --git a/src/sindi/comparator.py b/src/sindi/comparator.py index 81865c3..ae15cd9 100644 --- a/src/sindi/comparator.py +++ b/src/sindi/comparator.py @@ -6,6 +6,7 @@ from .simplifier import Simplifier from .rewriter import Rewriter from .utils import printer +from .ast_rewriter import ASTRewriter import z3 import re @@ -16,16 +17,34 @@ def __init__(self): self.simplifier = Simplifier() self.parser = Parser([]) self.rewriter = Rewriter() + self.ast_rewriter = ASTRewriter() + + # Old version. Keeping it for reference. + # def _parse_predicate(self, predicate_str: str) -> ASTNode: + # predicate_str = self.rewriter.apply(predicate_str) + # self.parser.tokens = self.tokenizer.tokenize(predicate_str) + # self.parser.pos = 0 + # return self.parser.parse() - def _parse_predicate(self, predicate_str): - predicate_str = self.rewriter.apply(predicate_str) - self.parser.tokens = self.tokenizer.tokenize(predicate_str) - self.parser.pos = 0 - return self.parser.parse() + def _parse_predicate(self, predicate_str: str) -> ASTNode: + """ + Single source of truth for: string rewrite -> tokenize -> parse -> AST normalize. + Keeping this here guarantees all compare paths see identical canonicalization. + """ + s = self.rewriter.apply(predicate_str) + tokens = self.tokenizer.tokenize(s) + ast = Parser(tokens).parse() + # AST-level normalization (boolean ==/!= to True/False, !!, move '-' across rels, sort, etc.) + try: + ast = self.ast_rewriter.normalize(ast) + except Exception: + # Never block compare() if AST-normalization adds a corner case later. + pass + return ast def compare(self, predicate1: str, predicate2: str) -> str: - predicate1 = self.rewriter.apply(predicate1) - predicate2 = self.rewriter.apply(predicate2) + # predicate1 = self.rewriter.apply(predicate1) + # predicate2 = self.rewriter.apply(predicate2) # Tokenize, parse, and simplify the first predicate tokens1 = self.tokenizer.tokenize(predicate1) @@ -41,6 +60,12 @@ def compare(self, predicate1: str, predicate2: str) -> str: ast2 = parser2.parse() printer(f"Parsed AST2: {ast2}") + # Parse both via the unified pipeline so string rewrites are always applied. + ast1 = self._parse_predicate(predicate1) + printer(f"Parsed+Normalized AST1: {ast1}") + ast2 = self._parse_predicate(predicate2) + printer(f"Parsed+Normalized AST2: {ast2}") + # Special-case: identical LHS/RHS with strict compare vs '!=' (both UNSAT), # but tests expect the '!=' side to be considered stronger. if self._is_strict_vs_neq_same_operands(ast1, ast2): diff --git a/src/sindi/rewriter.py b/src/sindi/rewriter.py index 064067e..b3d624a 100644 --- a/src/sindi/rewriter.py +++ b/src/sindi/rewriter.py @@ -5,6 +5,9 @@ class Rewriter: """ Canonicalizes Solidity predicate strings before tokenization/parsing, implementing the rewrite rules in Table \\ref{tab:canonicalized}. + This is a surface-level rewriting pass, working on the raw string which is + then tokenized and parsed into an AST. We have replaced this by the AST-level + rewriter class `ASTRewriter`. """ _HEX_ZERO_ADDR = re.compile(r"\b0x0{40}\b", flags=re.IGNORECASE) diff --git a/src/sindi/tokenizer.py b/src/sindi/tokenizer.py index 104b3ba..4f4463e 100644 --- a/src/sindi/tokenizer.py +++ b/src/sindi/tokenizer.py @@ -33,6 +33,9 @@ def __init__(self): (r'\]', 'RBRACKET'), (r'\"[^\"]*\"', 'STRING_LITERAL'), + # --- Handle 10**k wei as one numeric token (fallback if string rewriter didn't run) --- + (r'(?i)\b10\s*\*\*\s*(\d+)\s*wei\b', 'WEI_POW10'), + # ---- Numbers (order matters: scientific before float/int) ---- (r'\b\d(?:_?\d)*(?:\.\d(?:_?\d)*)?[eE][+-]?\d+(?:_?\d)*\b', 'SCIENTIFIC'), (r'\b\d(?:_?\d)*\.\d(?:_?\d)*\b', 'FLOAT'), @@ -83,6 +86,13 @@ def tokenize(self, predicate: str) -> List[Tuple[str, str]]: value = str(num * self.time_units[unit]) tag = 'INTEGER' + elif tag == 'WEI_POW10': + # Turn "10**18 wei" into a big integer literal token + m = re.match(r'(?i)\b10\s*\*\*\s*(\d+)\s*wei\b', value) + k = int(m.group(1)) if m else 0 + value = str(10 ** k) + tag = 'INTEGER' + elif tag in ('SCIENTIFIC', 'FLOAT', 'INTEGER'): # Strip underscores from numeric tokens value = value.replace('_', '')