122 lines
3.0 KiB
Python
122 lines
3.0 KiB
Python
# Example: expression tree + evaluation
|
|
|
|
from enum import Enum
|
|
|
|
|
|
class TokenType(Enum):
|
|
NUMBER = 1
|
|
OPERATOR = 2
|
|
|
|
|
|
# either store an operator (+, -, *, /) or a number
|
|
class Token:
|
|
def __init__(self, type: TokenType, value: str):
|
|
self.type = type
|
|
self.value = value
|
|
|
|
|
|
# each individual element in the tree
|
|
class Node:
|
|
left = None
|
|
right = None
|
|
|
|
def __init__(self, token: Token, left=None, right=None):
|
|
self.token = token
|
|
self.left = left
|
|
self.right = right
|
|
|
|
|
|
# holder for the tree
|
|
class ExpressionTree:
|
|
def __init__(self, root: Node | None):
|
|
self.root = root
|
|
|
|
|
|
def evaluate(tree: ExpressionTree) -> int:
|
|
return _evaluate(tree.root)
|
|
|
|
|
|
# evaluate a tree recursively starting from the node
|
|
def _evaluate(node: Node | None) -> int:
|
|
if node is None:
|
|
return 0
|
|
|
|
if node.token.type == TokenType.NUMBER:
|
|
return int(node.token.value)
|
|
|
|
elif node.token.type == TokenType.OPERATOR:
|
|
if node.left is None or node.right is None:
|
|
raise ValueError(f"Operator {node.token.value} requires two operands")
|
|
|
|
left = _evaluate(node.left)
|
|
right = _evaluate(node.right)
|
|
|
|
if node.token.value == "+":
|
|
return left + right
|
|
elif node.token.value == "-":
|
|
return left - right
|
|
elif node.token.value == "*":
|
|
return left * right
|
|
elif node.token.value == "/":
|
|
return left // right
|
|
|
|
raise ValueError(f"Unknown operator: {node.token.value}")
|
|
|
|
|
|
# test cases
|
|
test_cases = [
|
|
(ExpressionTree(Node(Token(TokenType.NUMBER, "5"))), 5),
|
|
(
|
|
ExpressionTree(
|
|
Node(
|
|
Token(TokenType.OPERATOR, "+"),
|
|
Node(Token(TokenType.NUMBER, "3")),
|
|
Node(Token(TokenType.NUMBER, "2")),
|
|
)
|
|
),
|
|
5,
|
|
),
|
|
(
|
|
ExpressionTree(
|
|
Node(
|
|
Token(TokenType.OPERATOR, "+"),
|
|
Node(Token(TokenType.NUMBER, "3")),
|
|
Node(
|
|
Token(TokenType.OPERATOR, "*"),
|
|
Node(Token(TokenType.NUMBER, "2")),
|
|
Node(Token(TokenType.NUMBER, "3")),
|
|
),
|
|
)
|
|
),
|
|
9,
|
|
),
|
|
(
|
|
ExpressionTree(
|
|
Node(
|
|
Token(TokenType.OPERATOR, "+"),
|
|
Node(
|
|
Token(TokenType.OPERATOR, "-"),
|
|
Node(Token(TokenType.NUMBER, "3")),
|
|
Node(Token(TokenType.NUMBER, "2")),
|
|
),
|
|
Node(
|
|
Token(TokenType.OPERATOR, "*"),
|
|
Node(Token(TokenType.NUMBER, "2")),
|
|
Node(Token(TokenType.NUMBER, "3")),
|
|
),
|
|
)
|
|
),
|
|
7,
|
|
),
|
|
]
|
|
|
|
|
|
def test_evaluate() -> None:
|
|
for tree, expected in test_cases:
|
|
evaluation = evaluate(tree) == expected
|
|
print(f"Test case: evaluate({tree}) == {expected} -> {evaluation}")
|
|
|
|
|
|
if __name__ == "__main__":
|
|
test_evaluate()
|