@ruleset
def math_ruleset(a: Math, b: Math, c: Math, i: i64, j: i64, xs: MultiSet[Math], ys: MultiSet[Math], zs: MultiSet[Math]):
yield rewrite(a + b).to(sum(MultiSet(a, b)))
yield rewrite(a * b).to(product(MultiSet(a, b)))
# 0 or 1 elements sums/products also can be extracted back to numbers
yield rule(a == sum(xs), xs.length() == i64(1)).then(a == xs.pick())
yield rule(a == product(xs), xs.length() == i64(1)).then(a == xs.pick())
yield rewrite(sum(MultiSet[Math]())).to(Math(0))
yield rewrite(product(MultiSet[Math]())).to(Math(1))
# distributive rule (a * (b + c) = a*b + a*c)
yield rule(
b == product(ys),
a == sum(xs),
ys.contains(a),
ys.length() > 1,
zs == ys.remove(a),
).then(
b == sum(xs.map(lambda x: product(zs.insert(x)))),
)
# constants
yield rule(
a == sum(xs),
b == Math(i),
xs.contains(b),
ys == xs.remove(b),
c == Math(j),
ys.contains(c),
).then(
a == sum(ys.remove(c).insert(Math(i + j))),
)
yield rule(
a == product(xs),
b == Math(i),
xs.contains(b),
ys == xs.remove(b),
c == Math(j),
ys.contains(c),
).then(
a == product(ys.remove(c).insert(Math(i * j))),
)