lib/chant/src/lower/expression/binary.zig

daab053ee43316e1809a84551d573ddd1e5bf3d2

 1 const choir = @import("choir");
 2 const ast = @import("../../ast/root.zig");
 3 const lower_mod = @import("../root.zig");
 4 const types = @import("types.zig");
 5 
 6 const ArithDialect = choir.dialects.ArithDialect;
 7 const CmpPredicate = choir.dialects.arith.CmpPredicate;
 8 const Error = lower_mod.Error;
 9 const Lowerer = lower_mod.Lowerer;
10 const Typed = types.Typed;
11 
12 const Pair = struct {
13     lhs: *choir.Value,
14     rhs: *choir.Value,
15     c_type: *const ast.Type,
16 };
17 
18 pub fn lower(lowerer: *Lowerer, binary: ast.expr.Binary, comptime recurse: types.LowerExpression) Error!Typed {
19     switch (binary.op) {
20         .logical_and, .logical_or => return error.UnsupportedConstruct,
21         .lt, .gt, .le, .ge, .eq, .ne => {
22             const operands = try lowerArithmeticPair(lowerer, binary.lhs, binary.rhs, recurse);
23             const predicate = comparisonPredicate(binary.op, operands.c_type);
24             const cmp = ArithDialect.CmpOp.create(lowerer.ctx, lowerer.loc, predicate, operands.lhs, operands.rhs) catch return error.OutOfMemory;
25             try lower_mod.emit.append(lowerer, cmp.op);
26             var mutable = cmp;
27             return .{ .value = mutable.getResult(), .c_type = &ast.types.int_type };
28         },
29         else => {},
30     }
31 
32     const operands = try lowerArithmeticPair(lowerer, binary.lhs, binary.rhs, recurse);
33     const op: *choir.Operation = switch (binary.op) {
34         .add => (ArithDialect.AddOp.create(lowerer.ctx, lowerer.loc, operands.lhs, operands.rhs) catch return error.OutOfMemory).op,
35         .sub => (ArithDialect.SubOp.create(lowerer.ctx, lowerer.loc, operands.lhs, operands.rhs) catch return error.OutOfMemory).op,
36         .mul => (ArithDialect.MulOp.create(lowerer.ctx, lowerer.loc, operands.lhs, operands.rhs) catch return error.OutOfMemory).op,
37         .div => (ArithDialect.DivOp.create(lowerer.ctx, lowerer.loc, operands.lhs, operands.rhs) catch return error.OutOfMemory).op,
38         .rem => (ArithDialect.RemOp.create(lowerer.ctx, lowerer.loc, operands.lhs, operands.rhs) catch return error.OutOfMemory).op,
39         .bit_and => (ArithDialect.AndOp.create(lowerer.ctx, lowerer.loc, operands.lhs, operands.rhs) catch return error.OutOfMemory).op,
40         .bit_or => (ArithDialect.OrOp.create(lowerer.ctx, lowerer.loc, operands.lhs, operands.rhs) catch return error.OutOfMemory).op,
41         .bit_xor => (ArithDialect.XorOp.create(lowerer.ctx, lowerer.loc, operands.lhs, operands.rhs) catch return error.OutOfMemory).op,
42         .shl => (ArithDialect.ShlOp.create(lowerer.ctx, lowerer.loc, operands.lhs, operands.rhs) catch return error.OutOfMemory).op,
43         .shr => (ArithDialect.ShrOp.create(lowerer.ctx, lowerer.loc, operands.lhs, operands.rhs) catch return error.OutOfMemory).op,
44         else => return error.UnsupportedConstruct,
45     };
46     try lower_mod.emit.append(lowerer, op);
47     const result = op.getResult(0) orelse return error.UnsupportedConstruct;
48     return .{ .value = result, .c_type = operands.c_type };
49 }
50 
51 fn lowerArithmeticPair(lowerer: *Lowerer, lhs_expr: *ast.Expr, rhs_expr: *ast.Expr, comptime recurse: types.LowerExpression) Error!Pair {
52     const lhs = try recurse(lowerer, lhs_expr);
53     const rhs = try recurse(lowerer, rhs_expr);
54     if (!ast.types.isArithmetic(lhs.c_type) or !ast.types.isArithmetic(rhs.c_type)) return error.UnsupportedConstruct;
55     const lhs_promoted = lower_mod.convert.promote(lhs.c_type);
56     const rhs_promoted = lower_mod.convert.promote(rhs.c_type);
57     const common = ast.types.commonArithmetic(lhs_promoted, rhs_promoted);
58     const lhs_value = try lower_mod.convert.convert(lowerer, lhs.value, lhs.c_type, common);
59     const rhs_value = try lower_mod.convert.convert(lowerer, rhs.value, rhs.c_type, common);
60     return .{ .lhs = lhs_value, .rhs = rhs_value, .c_type = common };
61 }
62 
63 fn comparisonPredicate(op: ast.expr.BinaryOp, c_type: *const ast.Type) CmpPredicate {
64     const unsigned = ast.types.isInteger(c_type) and c_type.is_unsigned;
65     return switch (op) {
66         .eq => .eq,
67         .ne => .ne,
68         .lt => if (unsigned) .ult else .lt,
69         .le => if (unsigned) .ule else .le,
70         .gt => if (unsigned) .ugt else .gt,
71         .ge => if (unsigned) .uge else .ge,
72         else => unreachable,
73     };
74 }