7. Kaleidoscope: Extending the Language: Mutable Variables

7.1 Chapter 7 Introduction

Welcome to Chapter 7 of the "Implementing a language with MLIR" tutorial. In chapters 1 through 6, we've built a very respectable, albeit simple, functional programming language. In our journey, we learned some parsing techniques, how to build and represent an AST, how to build MLIR, lower it to LLVM IR, optimize it, and JIT compile it.

While Kaleidoscope is interesting as a functional language, the fact that it is functional makes it "too easy" to generate SSA-based IR for it. In particular, a functional language makes it very easy to build values directly in SSA form. MLIR also represents values in SSA form, so this is a very nice property and it is often unclear to newcomers how to generate code for an imperative language with mutable variables.

The short (and happy) summary of this chapter is that there is no need for your front-end to build SSA form: MLIR provides reusable and well-tested support for this, though the way it works is a bit unexpected for some.

7.2 Why is this a hard problem?

To understand why mutable variables cause complexities in SSA construction, consider this extremely simple C example:

int G, H;
int test(_Bool Condition) {
  int X;
  if (Condition)
    X = G;
  else
    X = H;
  return X;
}

In this case, we have the variable "X", whose value depends on the path executed in the program. Because there are two different possible values for X before the return instruction, a PHI node is inserted to merge the two values. The LLVM IR that we want for this example looks like this:

@G = weak global i32 0   ; type of @G is i32*
@H = weak global i32 0   ; type of @H is i32*

define i32 @test(i1 %Condition) {
entry:
  br i1 %Condition, label %cond_true, label %cond_false

cond_true:
  %X.0 = load i32, i32* @G
  br label %cond_next

cond_false:
  %X.1 = load i32, i32* @H
  br label %cond_next

cond_next:
  %X.2 = phi i32 [ %X.1, %cond_false ], [ %X.0, %cond_true ]
  ret i32 %X.2
}

In this example, the loads from the G and H global variables are explicit in the LLVM IR, and they live in the then/else branches of the if statement (cond_true/cond_false). In order to merge the incoming values, the X.2 phi node in the cond_next block selects the right value to use based on where control flow is coming from: if control flow comes from the cond_false block, X.2 gets the value of X.1. Alternatively, if control flow comes from cond_true, it gets the value of X.0. The intent of this chapter is not to explain the details of SSA form. For more information, see one of the many online references.

The question for this chapter is: who places the MLIR block arguments when lowering assignments to mutable variables? As we saw in Chapter 5, we can solve this problem entirely in MLIR. The issue here is that MLIR represents values in SSA form: there is no "non-SSA" mode for values. However, SSA construction requires non-trivial algorithms and data structures, so it is inconvenient and wasteful for every frontend to have to reproduce this logic.

7.3 Memory in MLIR

The "trick" here is that while MLIR does require all values to be in SSA form, it does not require the contents of memory to be in SSA form. In the example above, note that the loads from G and H are direct accesses to G and H: the memory objects are not renamed or versioned. This differs from some other compiler systems, which do try to version memory objects. In MLIR, instead of encoding dataflow analysis of memory into the IR, it is handled with analyses and passes that are computed when needed.

With this in mind, the high-level idea is that we want to make a stack variable (which lives in memory, because it is on the stack) for each mutable object in a function. To take advantage of this trick, we need to talk about how MLIR represents stack variables.

In MLIR, all memory accesses are explicit with memref.load and memref.store operations, and there is no need for an "address-of" operator at this level. Notice how the type of %x in the example below is actually memref<f64> even though the variable holds an f64. What this means is that %x identifies space for an f64, and its name refers to that storage rather than to the value currently stored there. Stack variables are declared with the memref.alloca operation:

%x = memref.alloca() : memref<f64>
memref.store %initial, %x[] : memref<f64>
%value = memref.load %x[] : memref<f64>

This code shows an example of how you can declare and manipulate a stack variable in MLIR. Stack memory allocated with memref.alloca is fully general: you can pass the memref to functions, create aliases or views of it, and use it with other memory operations. In our example above, we could rewrite the example to use the memref.alloca technique to avoid creating block arguments in the frontend:

// example.mlir
module {
  func.func @test(%condition: i1, %g: f64, %h: f64) -> f64 {
    %x = memref.alloca() : memref<f64>
    scf.if %condition {
      memref.store %g, %x[] : memref<f64>
    } else {
      memref.store %h, %x[] : memref<f64>
    }
    %value = memref.load %x[] : memref<f64>
    return %value : f64
  }
}

If we lower this directly to LLVM IR without running mem2reg, the stack slot and its accesses remain. The insertvalue and extractvalue instructions construct and access the descriptor used while lowering the memref, but the important operations here are the alloca, the store in each branch, and the load after the branches merge:

define double @test(i1 %0, double %1, double %2) {
  %4 = alloca double, i64 1, align 8
  %5 = insertvalue { ptr, ptr, i64 } poison, ptr %4, 0
  %6 = insertvalue { ptr, ptr, i64 } %5, ptr %4, 1
  %7 = insertvalue { ptr, ptr, i64 } %6, i64 0, 2
  br i1 %0, label %8, label %10

8:
  %9 = extractvalue { ptr, ptr, i64 } %7, 1
  store double %1, ptr %9, align 8
  br label %12

10:
  %11 = extractvalue { ptr, ptr, i64 } %7, 1
  store double %2, ptr %11, align 8
  br label %12

12:
  %13 = extractvalue { ptr, ptr, i64 } %7, 1
  %14 = load double, ptr %13, align 8
  ret double %14
}

With this, we have discovered a way to handle arbitrary mutable variables without the need to create block arguments at all:

  1. Each mutable variable becomes a stack allocation.
  2. Each read of the variable becomes a load from the stack.
  3. Each update of the variable becomes a store to the stack.
  4. Passing a variable's storage uses the memref directly.

While this solution has solved our immediate problem, it introduced another one: we have now apparently introduced a lot of stack traffic for very simple and common operations, a major performance problem. Fortunately for us, MLIR has a mem2reg pass that handles this case, promoting suitable memory slots into SSA values and inserting block arguments as appropriate.

The mem2reg pass operates on control-flow blocks, so we first lower scf.if to the cf dialect. If you run this example through the passes, you'll get:

mlir-opt example.mlir --convert-scf-to-cf --mem2reg
module {
  func.func @test(%arg0: i1, %arg1: f64, %arg2: f64) -> f64 {
    cf.cond_br %arg0, ^bb1, ^bb2
  ^bb1:
    cf.br ^bb3(%arg1 : f64)
  ^bb2:
    cf.br ^bb3(%arg2 : f64)
  ^bb3(%0: f64):
    return %0 : f64
  }
}

The mem2reg pass implements the standard "iterated dominance frontier" algorithm for constructing SSA form. When multiple stored values can reach a load, it inserts a block argument at the merge point and passes the appropriate value from each predecessor. The mem2reg optimization pass is the answer to dealing with mutable variables, and we highly recommend that you depend on it. Note that MLIR's mem2reg only works on variables in certain circumstances:

  1. mem2reg is allocation-driven: it looks for operations that expose promotable memory slots. For this tutorial, those are memref.alloca operations. It does not promote global variables or arbitrary heap allocations.
  2. The allocation must dominate all of its uses. Creating our stack slots in the function's entry block makes them available throughout the function and makes this condition easy to satisfy.
  3. Every use of the memory slot must be understood by the promotion interfaces. Direct memref.load and memref.store operations work. Passing the memref to an arbitrary function or using it through an unsupported alias prevents promotion.
  4. A memref.alloca containing one element can be promoted directly to a scalar SSA value. This is why Kaleidoscope uses the zero-dimensional memref<f64> type for each mutable variable. MLIR can handle some more complicated memrefs too, but those cases are not needed here.

All of these properties are easy to satisfy for most imperative languages, and we'll illustrate it below with Kaleidoscope. The final question you may be asking is: should I bother with this nonsense for my frontend? Wouldn't it be better if I just did SSA construction directly, avoiding use of the mem2reg optimization pass? In short, we strongly recommend that you use this technique for building SSA form, unless there is an extremely good reason not to. Using this technique is:

  • Proven and well tested: MLIR provides the shared SSA-construction algorithm and promotion interfaces, so each frontend does not have to implement and maintain its own version.
  • Extremely fast: mem2reg forwards stored values directly to their uses and only adds block arguments at the merge points where they are needed.
  • Needed for debug info generation: keeping a variable in memory gives it a concrete address to which debug information can be attached. In Chapter 9, we retain that storage in unoptimized builds so the debugger can show our variables.

If nothing else, this makes it much easier to get your frontend up and running, and is very simple to implement. Let's extend Kaleidoscope with mutable variables now!

7.4 Mutable Variables in Kaleidoscope

Now that we know the sort of problem we want to tackle, let's see what this looks like in the context of our little Kaleidoscope language. We're going to add two features:

  1. The ability to mutate variables with the '=' operator.
  2. The ability to define new variables.

While the first item is really what this is about, we only have variables for incoming arguments as well as for induction variables, and redefining those only goes so far :). Also, the ability to define new variables is a useful thing regardless of whether you will be mutating them. Here's a motivating example that shows how we could use these:

# Define ':' for sequencing: as a low-precedence operator that ignores operands
# and just returns the RHS.
def binary : 1 (x y) y;

# Recursive fib, we could do this before.
def fib(x)
  if (x < 3) then
    1
  else
    fib(x-1)+fib(x-2);

# Iterative fib.
def fibi(x)
  var a = 1, b = 1, c in
  (for i = 3, i < x + 1 in
     c = a + b :
     a = b :
     b = c) :
  b;

# Call it.
fibi(10);

In order to mutate variables, we will change existing variables to use MLIR memory operations. Once that works, we will add assignment and then extend Kaleidoscope with local variable definitions.

7.5 Adjusting Existing Variables for Mutation

The symbol table in Kaleidoscope is managed at code generation time by the NamedValues map. This map currently keeps track of the MLIR Value that holds the f64 value for the named variable. In order to support mutation, we need to change this slightly, so that NamedValues holds the memory location of the variable in question. Note that this change is a refactoring: it changes the structure of the code, but does not (by itself) change the behavior of the compiler. All of these changes are isolated in the Kaleidoscope code generator.

At this point in Kaleidoscope's development, it only supports variables for two things: incoming arguments to functions and the induction variable of for loops. For consistency, we'll allow mutation of these variables in addition to other user-defined variables. This means that these will both need memory locations.

To start our transformation of Kaleidoscope, we'll change the meaning of the values stored in the NamedValues map. Instead of mapping each name directly to its current f64 SSA value, it will map the name to a zero-dimensional memref containing the variable. Both are represented by MLIR's Value C++ type, so the declaration itself does not change:

static std::map<std::string, Value> NamedValues;

We first find the function surrounding the builder's current insertion point, then use a helper to create storage at the beginning of that function:

static func::FuncOp getCurrentFunction() {
  Operation *Parent = TheBuilder->getInsertionBlock()->getParentOp();
  if (auto Function = dyn_cast<func::FuncOp>(Parent))
    return Function;
  return Parent->getParentOfType<func::FuncOp>();
}

/// CreateEntryBlockStorage - Create mutable storage in the function entry block.
static Value CreateEntryBlockStorage() {
  func::FuncOp Function = getCurrentFunction();
  OpBuilder::InsertionGuard Guard(*TheBuilder);
  TheBuilder->setInsertionPointToStart(&Function.front());
  auto VariableType = MemRefType::get({}, TheBuilder->getF64Type());
  return TheBuilder->create<memref::AllocaOp>(getLocation(), VariableType);
}

InsertionGuard restores the builder's previous insertion point when the helper returns. Placing the allocation at the start of the entry block makes the storage available to every region in the function.

A variable reference now looks up its storage and loads the current value:

Value VariableExprAST::codegen() {
  auto It = NamedValues.find(Name);
  if (It == NamedValues.end())
    return LogErrorV("Unknown variable name");

  return TheBuilder->create<memref::LoadOp>(getLocation(), It->second,
                                             ValueRange{});
}

Loop induction variables use the same representation. We allocate a slot, store the initial value, and make the slot visible through NamedValues:

Value Variable = CreateEntryBlockStorage();
TheBuilder->create<memref::StoreOp>(getLocation(), StartVal, Variable,
                                    ValueRange{});

auto OldValue = NamedValues.find(VarName);
bool HadOldValue = OldValue != NamedValues.end();
Value SavedValue = HadOldValue ? OldValue->second : Value();
NamedValues[VarName] = Variable;

After the loop body and step have run, we reload the induction variable in case either expression changed it, compute the next value, and store it:

Value Current =
    Builder.create<memref::LoadOp>(Loc, Variable, ValueRange{});
Value NextVar = Builder.create<arith::AddFOp>(Loc, Current, StepVal);
Builder.create<memref::StoreOp>(Loc, NextVar, Variable, ValueRange{});

Function arguments also become mutable storage when a function body is created:

NamedValues.clear();
unsigned Index = 0;
for (BlockArgument Argument : TheFunction.getArguments()) {
  Value Storage = CreateEntryBlockStorage();
  TheBuilder->create<memref::StoreOp>(getLocation(), Argument, Storage,
                                      ValueRange{});
  NamedValues[P.getArgs()[Index++]] = Storage;
}

Finally, the lowering pipeline converts structured control flow to blocks and runs MLIR's mem2reg pass before lowering the remaining operations to the LLVM dialect:

PassManager LoweringPM(TheContext.get());
LoweringPM.addPass(createSCFToControlFlowPass());
LoweringPM.addPass(createMem2Reg());
LoweringPM.addPass(createConvertFuncToLLVMPass());
LoweringPM.addPass(createArithToLLVMConversionPass());
LoweringPM.addPass(createFinalizeMemRefToLLVMConversionPass());
LoweringPM.addPass(createConvertControlFlowToLLVMPass());
LoweringPM.addPass(createReconcileUnrealizedCastsPass());

Now all variable references consistently go through mutable storage, while the promotion pass can recover SSA values before translation to LLVM IR.

7.6 New Assignment Operator

With our current framework, adding an assignment operator is simple. We parse it like any other binary operator, but handle it internally instead of allowing the user to define it. First we give it a low precedence:

int main() {
  // Install standard binary operators.
  // 1 is lowest precedence.
  BinopPrecedence['='] = 2;
  BinopPrecedence['<'] = 10;
  BinopPrecedence['+'] = 20;
  BinopPrecedence['-'] = 20;

Assignment cannot follow the usual “emit the LHS, emit the RHS, perform the operation” pattern because evaluating the LHS would load its value. Instead, we ask whether the left-hand expression names a variable, generate the new value, and store it into that variable's memref:

Value BinaryExprAST::codegen() {
  // Assignment stores into the variable's mutable memref slot.
  if (Op == '=') {
    const std::string *Name = LHS->getVariableName();
    if (!Name)
      return LogErrorV("destination of '=' must be a variable");

    Value AssignedValue = RHS->codegen();
    if (!AssignedValue)
      return {};

    auto It = NamedValues.find(*Name);
    if (It == NamedValues.end())
      return LogErrorV("Unknown variable name");

    TheBuilder->create<memref::StoreOp>(getLocation(), AssignedValue,
                                        It->second, ValueRange{});
    return AssignedValue;
  }

Returning the assigned value permits chained assignments such as x = (y = z). We can now mutate loop variables and function arguments:

# Function to print a double.
extern printd(x);

# Define ':' for sequencing.
def binary : 1 (x y) y;

def test(x)
  printd(x) :
  x = 4 :
  printd(x);

test(123);

When run, this prints 123 and then 4. We can mutate existing variables; next we will add declarations for new local variables.

7.7 User-defined Local Variables

Adding var/in is just like any other extension we made to Kaleidoscope: we extend the lexer, the parser, the AST and the code generator. The first step for adding our new 'var/in' construct is to extend the lexer. As before, this is pretty trivial, the code looks like this:

enum Token {
  ...
  // var definition
  tok_var = -13
...
}
...
static int gettok() {
...
    if (IdentifierStr == "in")
      return tok_in;
    if (IdentifierStr == "binary")
      return tok_binary;
    if (IdentifierStr == "unary")
      return tok_unary;
    if (IdentifierStr == "var")
      return tok_var;
    return tok_identifier;
...

The next step is to define the AST node that we will construct. For var/in, it looks like this:

/// VarExprAST - Expression class for var/in
class VarExprAST : public ExprAST {
  std::vector<std::pair<std::string, std::unique_ptr<ExprAST>>> VarNames;
  std::unique_ptr<ExprAST> Body;

public:
  VarExprAST(
      std::vector<std::pair<std::string, std::unique_ptr<ExprAST>>> VarNames,
      std::unique_ptr<ExprAST> Body)
      : VarNames(std::move(VarNames)), Body(std::move(Body)) {}

  Value codegen() override;
};

var/in allows a list of names to be defined all at once, and each name can optionally have an initializer value. As such, we capture this information in the VarNames vector. Also, var/in has a body, this body is allowed to access the variables defined by the var/in.

With this in place, we can define the parser pieces. The first thing we do is add it as a primary expression:

/// primary
///   ::= identifierexpr
///   ::= numberexpr
///   ::= parenexpr
///   ::= ifexpr
///   ::= forexpr
///   ::= varexpr
static std::unique_ptr<ExprAST> ParsePrimary() {
  switch (CurTok) {
  default:
    return LogError("unknown token when expecting an expression");
  case tok_identifier:
    return ParseIdentifierExpr();
  case tok_number:
    return ParseNumberExpr();
  case '(':
    return ParseParenExpr();
  case tok_if:
    return ParseIfExpr();
  case tok_for:
    return ParseForExpr();
  case tok_var:
    return ParseVarExpr();
  }
}

Next we define ParseVarExpr:

/// varexpr ::= 'var' identifier ('=' expression)?
//                    (',' identifier ('=' expression)?)* 'in' expression
static std::unique_ptr<ExprAST> ParseVarExpr() {
  getNextToken(); // eat the var.

  std::vector<std::pair<std::string, std::unique_ptr<ExprAST>>> VarNames;
  if (CurTok != tok_identifier)
    return LogError("expected identifier after var");

The first part of this code parses the list of identifier/expr pairs into the local VarNames vector.

while (true) {
  std::string Name = IdentifierStr;
  getNextToken(); // eat identifier.

  std::unique_ptr<ExprAST> Init;
  if (CurTok == '=') {
    getNextToken(); // eat '='.
    Init = ParseExpression();
    if (!Init)
      return nullptr;
  }

  VarNames.emplace_back(Name, std::move(Init));

  if (CurTok != ',')
    break;
  getNextToken(); // eat ','.
  if (CurTok != tok_identifier)
    return LogError("expected identifier list after var");
}

Once all the variables are parsed, we then parse the body and create the AST node:

  if (CurTok != tok_in)
    return LogError("expected 'in' keyword after 'var'");
  getNextToken(); // eat 'in'.

  auto Body = ParseExpression();
  if (!Body)
    return nullptr;

  return std::make_unique<VarExprAST>(std::move(VarNames), std::move(Body));
}

Now that we can parse and represent the code, we generate mutable storage for each local variable:

Value VarExprAST::codegen() {
  std::vector<std::pair<std::string, std::optional<Value>>> OldBindings;

  auto RestoreBindings = [&]() {
    for (auto It = OldBindings.rbegin(); It != OldBindings.rend(); ++It) {
      if (It->second)
        NamedValues[It->first] = *It->second;
      else
        NamedValues.erase(It->first);
    }
  };

  for (auto &Variable : VarNames) {
    const std::string &Name = Variable.first;

    // Generate the initializer before introducing the new binding.
    Value InitialValue;
    if (Variable.second)
      InitialValue = Variable.second->codegen();
    else
      InitialValue = TheBuilder->create<arith::ConstantOp>(
          getLocation(), TheBuilder->getF64FloatAttr(0.0));
    if (!InitialValue) {
      RestoreBindings();
      return {};
    }

    Value Storage = CreateEntryBlockStorage();
    TheBuilder->create<memref::StoreOp>(getLocation(), InitialValue, Storage,
                                        ValueRange{});

    auto Old = NamedValues.find(Name);
    OldBindings.emplace_back(
        Name, Old == NamedValues.end() ? std::optional<Value>()
                                      : std::optional<Value>(Old->second));
    NamedValues[Name] = Storage;
  }

  Value BodyValue = Body->codegen();
  RestoreBindings();
  return BodyValue;
}

The initializer is generated before the new name is installed, so an inner declaration such as var a = a in ... reads the outer a. We remember every previous binding and restore them in reverse order when the body is finished.

With this, the iterative Fibonacci example from the introduction compiles and runs. The frontend expresses mutation with straightforward memory operations; MLIR's mem2reg pass promotes those operations into SSA values and introduces block arguments where different control-flow paths meet.

7.8 Download Source Code

https://github.com/alankarmisra/kaleidoscope-mlir-tutorial/

cd kaleidoscope-mlir-tutorial/code/chapter-07

7.9 Full Code Listing

Here is the complete code listing for our running example, enhanced with mutable variables and var/in support. Here is the CMake configuration:

cmake_minimum_required(VERSION 3.20)

project(kaleidoscope-chapter-07 LANGUAGES C CXX)

set(CMAKE_CXX_STANDARD 17)
set(CMAKE_CXX_STANDARD_REQUIRED YES)
set(CMAKE_CXX_EXTENSIONS NO)

find_package(MLIR REQUIRED CONFIG)

add_executable(toy toy.cpp)

# Make symbols in the executable available to the JIT for runtime lookup.
set_target_properties(toy PROPERTIES ENABLE_EXPORTS ON)

target_include_directories(toy PRIVATE
  ${LLVM_INCLUDE_DIRS}
  ${MLIR_INCLUDE_DIRS}
)

target_compile_definitions(toy PRIVATE ${LLVM_DEFINITIONS})

if(NOT LLVM_ENABLE_RTTI)
  target_compile_options(toy PRIVATE -fno-rtti)
endif()

llvm_map_components_to_libnames(LLVM_LIBS
  Core
  OrcJIT
  Support
  native
)

target_link_libraries(toy PRIVATE
  MLIRArithToLLVM
  MLIRArithDialect
  MLIRBuiltinToLLVMIRTranslation
  MLIRControlFlowDialect
  MLIRControlFlowToLLVM
  MLIRFuncToLLVM
  MLIRFuncDialect
  MLIRLLVMDialect
  MLIRLLVMToLLVMIRTranslation
  MLIRMemRefDialect
  MLIRMemRefToLLVM
  MLIRReconcileUnrealizedCasts
  MLIRSCFDialect
  MLIRSCFToControlFlow
  MLIRTargetLLVMIRExport
  MLIRTransforms
  ${LLVM_LIBS}
)

To build this example, use:

cmake -S . -B build \
  -DMLIR_DIR=/path/to/llvm-project/build/lib/cmake/mlir
cmake --build build
./build/toy

Here is the code:

#include "../include/KaleidoscopeJIT.h"
#include "mlir/Conversion/ArithToLLVM/ArithToLLVM.h"
#include "mlir/Conversion/ControlFlowToLLVM/ControlFlowToLLVM.h"
#include "mlir/Conversion/FuncToLLVM/ConvertFuncToLLVMPass.h"
#include "mlir/Conversion/MemRefToLLVM/MemRefToLLVM.h"
#include "mlir/Conversion/ReconcileUnrealizedCasts/ReconcileUnrealizedCasts.h"
#include "mlir/Conversion/SCFToControlFlow/SCFToControlFlow.h"
#include "mlir/Dialect/Arith/IR/Arith.h"
#include "mlir/Dialect/ControlFlow/IR/ControlFlowOps.h"
#include "mlir/Dialect/Func/IR/FuncOps.h"
#include "mlir/Dialect/MemRef/IR/MemRef.h"
#include "mlir/Dialect/SCF/IR/SCF.h"
#include "mlir/IR/Builders.h"
#include "mlir/IR/BuiltinOps.h"
#include "mlir/IR/MLIRContext.h"
#include "mlir/IR/OperationSupport.h"
#include "mlir/IR/Verifier.h"
#include "mlir/Pass/PassManager.h"
#include "mlir/Target/LLVMIR/Dialect/Builtin/BuiltinToLLVMIRTranslation.h"
#include "mlir/Target/LLVMIR/Dialect/LLVMIR/LLVMToLLVMIRTranslation.h"
#include "mlir/Target/LLVMIR/Export.h"
#include "mlir/Transforms/Passes.h"
#include "llvm/Support/CommandLine.h"
#include "llvm/Support/Error.h"
#include "llvm/Support/TargetSelect.h"
#include "llvm/Support/raw_ostream.h"
#include <cassert>
#include <cctype>
#include <cstdio>
#include <cstdlib>
#include <map>
#include <memory>
#include <optional>
#include <string>
#include <utility>
#include <vector>

using namespace mlir;

//===----------------------------------------------------------------------===//
// Lexer
//===----------------------------------------------------------------------===//

// The lexer returns tokens [0-255] if it is an unknown character, otherwise one
// of these for known things.
enum Token {
  tok_eof = -1,

  // commands
  tok_def = -2,
  tok_extern = -3,

  // primary
  tok_identifier = -4,
  tok_number = -5,

  // control
  tok_if = -6,
  tok_then = -7,
  tok_else = -8,
  tok_for = -9,
  tok_in = -10,

  // operators
  tok_binary = -11,
  tok_unary = -12,

  // var definition
  tok_var = -13
};

static std::string IdentifierStr; // Filled in for identifiers and keywords
static double NumVal;             // Filled in if tok_number

/// gettok - Return the next token from standard input.
static int gettok() {
  static int LastChar = ' ';

  // Skip any whitespace.
  while (isspace(LastChar))
    LastChar = getchar();

  if (isalpha(LastChar)) { // identifier: [a-zA-Z][a-zA-Z0-9]*
    IdentifierStr = LastChar;
    while (isalnum((LastChar = getchar())))
      IdentifierStr += LastChar;

    if (IdentifierStr == "def")
      return tok_def;
    if (IdentifierStr == "extern")
      return tok_extern;
    if (IdentifierStr == "if")
      return tok_if;
    if (IdentifierStr == "then")
      return tok_then;
    if (IdentifierStr == "else")
      return tok_else;
    if (IdentifierStr == "for")
      return tok_for;
    if (IdentifierStr == "in")
      return tok_in;
    if (IdentifierStr == "binary")
      return tok_binary;
    if (IdentifierStr == "unary")
      return tok_unary;
    if (IdentifierStr == "var")
      return tok_var;
    return tok_identifier;
  }

  if (isdigit(LastChar) || LastChar == '.') { // Number: [0-9.]+
    std::string NumStr;
    do {
      NumStr += LastChar;
      LastChar = getchar();
    } while (isdigit(LastChar) || LastChar == '.');

    NumVal = strtod(NumStr.c_str(), nullptr);
    return tok_number;
  }

  if (LastChar == '#') {
    // Comment until end of line.
    do
      LastChar = getchar();
    while (LastChar != EOF && LastChar != '\n' && LastChar != '\r');

    if (LastChar != EOF)
      return gettok();
  }

  // Check for end of file.  Don't eat the EOF.
  if (LastChar == EOF)
    return tok_eof;

  // Otherwise, just return the character as its ascii value.
  int ThisChar = LastChar;
  LastChar = getchar();
  return ThisChar;
}

//===----------------------------------------------------------------------===//
// Abstract Syntax Tree (aka Parse Tree)
//===----------------------------------------------------------------------===//

namespace {

/// ExprAST - Base class for all expression nodes.
class ExprAST {
public:
  virtual ~ExprAST() = default;

  virtual Value codegen() = 0;
  virtual const std::string *getVariableName() const { return nullptr; }
};

/// NumberExprAST - Expression class for numeric literals like "1.0".
class NumberExprAST : public ExprAST {
  double Val;

public:
  NumberExprAST(double Val) : Val(Val) {}

  Value codegen() override;
};

/// VariableExprAST - Expression class for referencing a variable, like "a".
class VariableExprAST : public ExprAST {
  std::string Name;

public:
  VariableExprAST(const std::string &Name) : Name(Name) {}

  Value codegen() override;
  const std::string *getVariableName() const override { return &Name; }
};

/// UnaryExprAST - Expression class for a unary operator.
class UnaryExprAST : public ExprAST {
  char Opcode;
  std::unique_ptr<ExprAST> Operand;

public:
  UnaryExprAST(char Opcode, std::unique_ptr<ExprAST> Operand)
      : Opcode(Opcode), Operand(std::move(Operand)) {}

  Value codegen() override;
};

/// BinaryExprAST - Expression class for a binary operator.
class BinaryExprAST : public ExprAST {
  char Op;
  std::unique_ptr<ExprAST> LHS, RHS;

public:
  BinaryExprAST(char Op, std::unique_ptr<ExprAST> LHS,
                std::unique_ptr<ExprAST> RHS)
      : Op(Op), LHS(std::move(LHS)), RHS(std::move(RHS)) {}

  Value codegen() override;
};

/// CallExprAST - Expression class for function calls.
class CallExprAST : public ExprAST {
  std::string Callee;
  std::vector<std::unique_ptr<ExprAST>> Args;

public:
  CallExprAST(const std::string &Callee,
              std::vector<std::unique_ptr<ExprAST>> Args)
      : Callee(Callee), Args(std::move(Args)) {}

  Value codegen() override;
};

/// IfExprAST - Expression class for if/then/else.
class IfExprAST : public ExprAST {
  std::unique_ptr<ExprAST> Cond, Then, Else;

public:
  IfExprAST(std::unique_ptr<ExprAST> Cond, std::unique_ptr<ExprAST> Then,
            std::unique_ptr<ExprAST> Else)
      : Cond(std::move(Cond)), Then(std::move(Then)), Else(std::move(Else)) {}

  Value codegen() override;
};

/// ForExprAST - Expression class for for/in.
class ForExprAST : public ExprAST {
  std::string VarName;
  std::unique_ptr<ExprAST> Start, End, Step, Body;

public:
  ForExprAST(const std::string &VarName, std::unique_ptr<ExprAST> Start,
             std::unique_ptr<ExprAST> End, std::unique_ptr<ExprAST> Step,
             std::unique_ptr<ExprAST> Body)
      : VarName(VarName), Start(std::move(Start)), End(std::move(End)),
        Step(std::move(Step)), Body(std::move(Body)) {}

  Value codegen() override;
};

/// VarExprAST - Expression class for var/in.
class VarExprAST : public ExprAST {
  std::vector<std::pair<std::string, std::unique_ptr<ExprAST>>> VarNames;
  std::unique_ptr<ExprAST> Body;

public:
  VarExprAST(
      std::vector<std::pair<std::string, std::unique_ptr<ExprAST>>> VarNames,
      std::unique_ptr<ExprAST> Body)
      : VarNames(std::move(VarNames)), Body(std::move(Body)) {}

  Value codegen() override;
};

/// PrototypeAST - This class represents the "prototype" for a function,
/// which captures its argument names as well as if it is an operator.
class PrototypeAST {
  std::string Name;
  std::vector<std::string> Args;
  bool IsOperator;
  unsigned Precedence; // Precedence if a binary op.

public:
  PrototypeAST(const std::string &Name, std::vector<std::string> Args,
               bool IsOperator = false, unsigned Prec = 0)
      : Name(Name), Args(std::move(Args)), IsOperator(IsOperator),
        Precedence(Prec) {}

  func::FuncOp codegen();
  const std::string &getName() const { return Name; }
  const std::vector<std::string> &getArgs() const { return Args; }

  bool isUnaryOp() const { return IsOperator && Args.size() == 1; }
  bool isBinaryOp() const { return IsOperator && Args.size() == 2; }

  char getOperatorName() const {
    assert(isUnaryOp() || isBinaryOp());
    return Name.back();
  }

  unsigned getBinaryPrecedence() const { return Precedence; }
};

/// FunctionAST - This class represents a function definition itself.
class FunctionAST {
  std::unique_ptr<PrototypeAST> Proto;
  std::unique_ptr<ExprAST> Body;

public:
  FunctionAST(std::unique_ptr<PrototypeAST> Proto,
              std::unique_ptr<ExprAST> Body)
      : Proto(std::move(Proto)), Body(std::move(Body)) {}

  func::FuncOp codegen();
};

} // end anonymous namespace

//===----------------------------------------------------------------------===//
// Parser
//===----------------------------------------------------------------------===//

/// CurTok/getNextToken - Provide a simple token buffer.  CurTok is the current
/// token the parser is looking at.  getNextToken reads another token from the
/// lexer and updates CurTok with its results.
static int CurTok;
static int getNextToken() { return CurTok = gettok(); }

/// BinopPrecedence - This holds the precedence for each binary operator that is
/// defined.
static std::map<char, int> BinopPrecedence;

/// GetTokPrecedence - Get the precedence of the pending binary operator token.
static int GetTokPrecedence() {
  if (!isascii(CurTok))
    return -1;

  // Make sure it's a declared binop.
  int TokPrec = BinopPrecedence[CurTok];
  if (TokPrec <= 0)
    return -1;
  return TokPrec;
}

/// LogError* - These are little helper functions for error handling.
std::unique_ptr<ExprAST> LogError(const char *Str) {
  fprintf(stderr, "Error: %s\n", Str);
  return nullptr;
}
std::unique_ptr<PrototypeAST> LogErrorP(const char *Str) {
  LogError(Str);
  return nullptr;
}

static std::unique_ptr<ExprAST> ParseExpression();

/// numberexpr ::= number
static std::unique_ptr<ExprAST> ParseNumberExpr() {
  auto Result = std::make_unique<NumberExprAST>(NumVal);
  getNextToken(); // consume the number
  return std::move(Result);
}

/// parenexpr ::= '(' expression ')'
static std::unique_ptr<ExprAST> ParseParenExpr() {
  getNextToken(); // eat (.
  auto V = ParseExpression();
  if (!V)
    return nullptr;

  if (CurTok != ')')
    return LogError("expected ')'");
  getNextToken(); // eat ).
  return V;
}

/// identifierexpr
///   ::= identifier
///   ::= identifier '(' expression* ')'
static std::unique_ptr<ExprAST> ParseIdentifierExpr() {
  std::string IdName = IdentifierStr;

  getNextToken(); // eat identifier.

  if (CurTok != '(') // Simple variable ref.
    return std::make_unique<VariableExprAST>(IdName);

  // Call.
  getNextToken(); // eat (
  std::vector<std::unique_ptr<ExprAST>> Args;
  if (CurTok != ')') {
    while (true) {
      if (auto Arg = ParseExpression())
        Args.push_back(std::move(Arg));
      else
        return nullptr;

      if (CurTok == ')')
        break;

      if (CurTok != ',')
        return LogError("Expected ')' or ',' in argument list");
      getNextToken();
    }
  }

  // Eat the ')'.
  getNextToken();

  return std::make_unique<CallExprAST>(IdName, std::move(Args));
}

/// ifexpr ::= 'if' expression 'then' expression 'else' expression
static std::unique_ptr<ExprAST> ParseIfExpr() {
  getNextToken(); // eat the if.

  auto Cond = ParseExpression();
  if (!Cond)
    return nullptr;

  if (CurTok != tok_then)
    return LogError("expected then");
  getNextToken(); // eat the then.

  auto Then = ParseExpression();
  if (!Then)
    return nullptr;

  if (CurTok != tok_else)
    return LogError("expected else");
  getNextToken(); // eat the else.

  auto Else = ParseExpression();
  if (!Else)
    return nullptr;

  return std::make_unique<IfExprAST>(std::move(Cond), std::move(Then),
                                     std::move(Else));
}

/// forexpr ::= 'for' identifier '=' expr ',' expr (',' expr)? 'in' expression
static std::unique_ptr<ExprAST> ParseForExpr() {
  getNextToken(); // eat the for.

  if (CurTok != tok_identifier)
    return LogError("expected identifier after for");

  std::string IdName = IdentifierStr;
  getNextToken(); // eat identifier.

  if (CurTok != '=')
    return LogError("expected '=' after for");
  getNextToken(); // eat '='.

  auto Start = ParseExpression();
  if (!Start)
    return nullptr;
  if (CurTok != ',')
    return LogError("expected ',' after for start value");
  getNextToken();

  auto End = ParseExpression();
  if (!End)
    return nullptr;

  // The step value is optional.
  std::unique_ptr<ExprAST> Step;
  if (CurTok == ',') {
    getNextToken();
    Step = ParseExpression();
    if (!Step)
      return nullptr;
  }

  if (CurTok != tok_in)
    return LogError("expected 'in' after for");
  getNextToken(); // eat the in.

  auto Body = ParseExpression();
  if (!Body)
    return nullptr;

  return std::make_unique<ForExprAST>(IdName, std::move(Start), std::move(End),
                                      std::move(Step), std::move(Body));
}

/// varexpr ::= 'var' identifier ('=' expression)?
///                    (',' identifier ('=' expression)?)* 'in' expression
static std::unique_ptr<ExprAST> ParseVarExpr() {
  getNextToken(); // eat the var.

  std::vector<std::pair<std::string, std::unique_ptr<ExprAST>>> VarNames;
  if (CurTok != tok_identifier)
    return LogError("expected identifier after var");

  while (true) {
    std::string Name = IdentifierStr;
    getNextToken(); // eat identifier.

    std::unique_ptr<ExprAST> Init;
    if (CurTok == '=') {
      getNextToken(); // eat '='.
      Init = ParseExpression();
      if (!Init)
        return nullptr;
    }

    VarNames.emplace_back(Name, std::move(Init));

    if (CurTok != ',')
      break;
    getNextToken(); // eat ','.
    if (CurTok != tok_identifier)
      return LogError("expected identifier list after var");
  }

  if (CurTok != tok_in)
    return LogError("expected 'in' keyword after 'var'");
  getNextToken(); // eat 'in'.

  auto Body = ParseExpression();
  if (!Body)
    return nullptr;

  return std::make_unique<VarExprAST>(std::move(VarNames), std::move(Body));
}

/// primary
///   ::= identifierexpr
///   ::= numberexpr
///   ::= parenexpr
///   ::= ifexpr
///   ::= forexpr
///   ::= varexpr
static std::unique_ptr<ExprAST> ParsePrimary() {
  switch (CurTok) {
  default:
    return LogError("unknown token when expecting an expression");
  case tok_identifier:
    return ParseIdentifierExpr();
  case tok_number:
    return ParseNumberExpr();
  case '(':
    return ParseParenExpr();
  case tok_if:
    return ParseIfExpr();
  case tok_for:
    return ParseForExpr();
  case tok_var:
    return ParseVarExpr();
  }
}

/// unary
///   ::= primary
///   ::= '!' unary
static std::unique_ptr<ExprAST> ParseUnary() {
  // If the current token is not an operator, it must be a primary expression.
  if (!isascii(CurTok) || CurTok == '(' || CurTok == ',')
    return ParsePrimary();

  // If this is a unary operator, read it.
  int Opc = CurTok;
  getNextToken();
  if (auto Operand = ParseUnary())
    return std::make_unique<UnaryExprAST>(Opc, std::move(Operand));
  return nullptr;
}

/// binoprhs
///   ::= ('+' unary)*
static std::unique_ptr<ExprAST> ParseBinOpRHS(int ExprPrec,
                                              std::unique_ptr<ExprAST> LHS) {
  // If this is a binop, find its precedence.
  while (true) {
    int TokPrec = GetTokPrecedence();

    // If this is a binop that binds at least as tightly as the current binop,
    // consume it, otherwise we are done.
    if (TokPrec < ExprPrec)
      return LHS;

    // Okay, we know this is a binop.
    int BinOp = CurTok;
    getNextToken(); // eat binop

    // Parse the unary expression after the binary operator.
    auto RHS = ParseUnary();
    if (!RHS)
      return nullptr;

    // If BinOp binds less tightly with RHS than the operator after RHS, let
    // the pending operator take RHS as its LHS.
    int NextPrec = GetTokPrecedence();
    if (TokPrec < NextPrec) {
      RHS = ParseBinOpRHS(TokPrec + 1, std::move(RHS));
      if (!RHS)
        return nullptr;
    }

    // Merge LHS/RHS.
    LHS =
        std::make_unique<BinaryExprAST>(BinOp, std::move(LHS), std::move(RHS));
  }
}

/// expression
///   ::= unary binoprhs
///
static std::unique_ptr<ExprAST> ParseExpression() {
  auto LHS = ParseUnary();
  if (!LHS)
    return nullptr;

  return ParseBinOpRHS(0, std::move(LHS));
}

/// prototype
///   ::= id '(' id* ')'
///   ::= binary LETTER number? (id, id)
///   ::= unary LETTER (id)
static std::unique_ptr<PrototypeAST> ParsePrototype() {
  std::string FnName;
  unsigned Kind = 0; // 0 = identifier, 1 = unary, 2 = binary.
  unsigned BinaryPrecedence = 30;

  switch (CurTok) {
  default:
    return LogErrorP("Expected function name in prototype");
  case tok_identifier:
    FnName = IdentifierStr;
    getNextToken();
    break;
  case tok_unary:
    getNextToken();
    if (!isascii(CurTok))
      return LogErrorP("Expected unary operator");
    FnName = "unary";
    FnName += static_cast<char>(CurTok);
    Kind = 1;
    getNextToken();
    break;
  case tok_binary:
    getNextToken();
    if (!isascii(CurTok))
      return LogErrorP("Expected binary operator");
    FnName = "binary";
    FnName += static_cast<char>(CurTok);
    Kind = 2;
    getNextToken();

    // Read the precedence if present.
    if (CurTok == tok_number) {
      if (NumVal < 1 || NumVal > 100)
        return LogErrorP("Invalid precedence: must be 1..100");
      BinaryPrecedence = static_cast<unsigned>(NumVal);
      getNextToken();
    }
    break;
  }

  if (CurTok != '(')
    return LogErrorP("Expected '(' in prototype");

  std::vector<std::string> ArgNames;
  while (getNextToken() == tok_identifier)
    ArgNames.push_back(IdentifierStr);
  if (CurTok != ')')
    return LogErrorP("Expected ')' in prototype");

  // success.
  getNextToken(); // eat ')'.

  // Verify the right number of names for an operator.
  if (Kind && ArgNames.size() != Kind)
    return LogErrorP("Invalid number of operands for operator");

  return std::make_unique<PrototypeAST>(FnName, std::move(ArgNames), Kind != 0,
                                        BinaryPrecedence);
}

/// definition ::= 'def' prototype expression
static std::unique_ptr<FunctionAST> ParseDefinition() {
  getNextToken(); // eat def.
  auto Proto = ParsePrototype();
  if (!Proto)
    return nullptr;

  if (auto E = ParseExpression())
    return std::make_unique<FunctionAST>(std::move(Proto), std::move(E));
  return nullptr;
}

/// toplevelexpr ::= expression
static std::unique_ptr<FunctionAST> ParseTopLevelExpr() {
  if (auto E = ParseExpression()) {
    // Make an anonymous proto.
    auto Proto = std::make_unique<PrototypeAST>("__anon_expr",
                                                std::vector<std::string>());
    return std::make_unique<FunctionAST>(std::move(Proto), std::move(E));
  }
  return nullptr;
}

/// external ::= 'extern' prototype
static std::unique_ptr<PrototypeAST> ParseExtern() {
  getNextToken(); // eat extern.
  return ParsePrototype();
}

//===----------------------------------------------------------------------===//
// Code Generation
//===----------------------------------------------------------------------===//

static std::unique_ptr<MLIRContext> TheContext;
static OwningOpRef<ModuleOp> TheModule;
static std::unique_ptr<OpBuilder> TheBuilder;
static std::unique_ptr<PassManager> ThePM;
static std::map<std::string, Value> NamedValues;
static std::unique_ptr<llvm::orc::KaleidoscopeJIT> TheJIT;
static std::map<std::string, std::unique_ptr<PrototypeAST>> FunctionProtos;
static llvm::ExitOnError ExitOnErr;
static llvm::cl::opt<bool> DumpMLIR(
    "dump-mlir", llvm::cl::desc("Print generated MLIR"),
    llvm::cl::init(false));
static llvm::cl::opt<bool>
    DumpLLVMIR("dump-llvm-ir",
             llvm::cl::desc("Print LLVM IR before adding it to the JIT"),
             llvm::cl::init(false));

static Location getLocation() { return TheBuilder->getUnknownLoc(); }

Value LogErrorV(const char *Str) {
  LogError(Str);
  return {};
}

func::FuncOp getFunction(const std::string &Name) {
  // First, see if the function has already been added to the current module.
  if (auto Function = TheModule->lookupSymbol<func::FuncOp>(Name))
    return Function;

  // If not, codegen the declaration from an existing prototype.
  auto It = FunctionProtos.find(Name);
  if (It != FunctionProtos.end()) {
    auto Function = It->second->codegen();
    Function.setPrivate();
    return Function;
  }

  return {};
}

static func::FuncOp getCurrentFunction() {
  Operation *Parent = TheBuilder->getInsertionBlock()->getParentOp();
  if (auto Function = dyn_cast<func::FuncOp>(Parent))
    return Function;
  return Parent->getParentOfType<func::FuncOp>();
}

/// CreateEntryBlockStorage - Create mutable storage in the function entry block.
static Value CreateEntryBlockStorage() {
  func::FuncOp Function = getCurrentFunction();
  OpBuilder::InsertionGuard Guard(*TheBuilder);
  TheBuilder->setInsertionPointToStart(&Function.front());
  auto VariableType = MemRefType::get({}, TheBuilder->getF64Type());
  return TheBuilder->create<memref::AllocaOp>(getLocation(), VariableType);
}

Value NumberExprAST::codegen() {
  return TheBuilder->create<arith::ConstantOp>(
      getLocation(), TheBuilder->getF64FloatAttr(Val));
}

Value VariableExprAST::codegen() {
  // Look this variable up in the function.
  auto It = NamedValues.find(Name);
  if (It == NamedValues.end())
    return LogErrorV("Unknown variable name");

  return TheBuilder->create<memref::LoadOp>(getLocation(), It->second,
                                             ValueRange{});
}

Value UnaryExprAST::codegen() {
  Value OperandV = Operand->codegen();
  if (!OperandV)
    return {};

  auto Operator = getFunction(std::string("unary") + Opcode);
  if (!Operator)
    return LogErrorV("Unknown unary operator");

  return TheBuilder->create<func::CallOp>(getLocation(), Operator, OperandV)
      .getResult(0);
}

Value BinaryExprAST::codegen() {
  // Assignment stores into the variable's mutable memref slot.
  if (Op == '=') {
    const std::string *Name = LHS->getVariableName();
    if (!Name)
      return LogErrorV("destination of '=' must be a variable");

    Value AssignedValue = RHS->codegen();
    if (!AssignedValue)
      return {};

    auto It = NamedValues.find(*Name);
    if (It == NamedValues.end())
      return LogErrorV("Unknown variable name");

    TheBuilder->create<memref::StoreOp>(getLocation(), AssignedValue,
                                        It->second, ValueRange{});
    return AssignedValue;
  }

  Value L = LHS->codegen();
  Value R = RHS->codegen();
  if (!L || !R)
    return {};

  switch (Op) {
  case '+':
    return TheBuilder->create<arith::AddFOp>(getLocation(), L, R);
  case '-':
    return TheBuilder->create<arith::SubFOp>(getLocation(), L, R);
  case '*':
    return TheBuilder->create<arith::MulFOp>(getLocation(), L, R);
  case '<': {
    Value Comparison = TheBuilder->create<arith::CmpFOp>(
        getLocation(), arith::CmpFPredicate::ULT, L, R);
    // Convert bool 0/1 to double 0.0 or 1.0.
    return TheBuilder->create<arith::UIToFPOp>(
        getLocation(), TheBuilder->getF64Type(), Comparison);
  }
  default:
    break;
  }

  // If it wasn't a builtin binary operator, it must be a user-defined one.
  auto Operator = getFunction(std::string("binary") + Op);
  if (!Operator)
    return LogErrorV("Unknown binary operator");

  Value Operands[] = {L, R};
  return TheBuilder->create<func::CallOp>(getLocation(), Operator, Operands)
      .getResult(0);
}

Value CallExprAST::codegen() {
  // Look up the name in the global module table.
  auto CalleeF = getFunction(Callee);
  if (!CalleeF)
    return LogErrorV("Unknown function referenced");

  // If argument mismatch error.
  if (CalleeF.getNumArguments() != Args.size())
    return LogErrorV("Incorrect # arguments passed");

  std::vector<Value> ArgsV;
  for (auto &Arg : Args) {
    ArgsV.push_back(Arg->codegen());
    if (!ArgsV.back())
      return {};
  }

  return TheBuilder->create<func::CallOp>(getLocation(), CalleeF, ArgsV)
      .getResult(0);
}

Value IfExprAST::codegen() {
  Value CondV = Cond->codegen();
  if (!CondV)
    return {};

  // Convert the condition to a boolean by comparing it with 0.0.
  Value Zero = TheBuilder->create<arith::ConstantOp>(
      getLocation(), TheBuilder->getF64FloatAttr(0.0));
  CondV = TheBuilder->create<arith::CmpFOp>(
      getLocation(), arith::CmpFPredicate::ONE, CondV, Zero);

  bool CodegenFailed = false;
  auto IfOp = TheBuilder->create<scf::IfOp>(
      getLocation(), CondV,
      [&](OpBuilder &Builder, Location Loc) {
        Value ThenV = Then->codegen();
        if (!ThenV) {
          CodegenFailed = true;
          ThenV = Builder.create<arith::ConstantOp>(
              Loc, Builder.getF64FloatAttr(0.0));
        }
        Builder.create<scf::YieldOp>(Loc, ThenV);
      },
      [&](OpBuilder &Builder, Location Loc) {
        Value ElseV = Else->codegen();
        if (!ElseV) {
          CodegenFailed = true;
          ElseV = Builder.create<arith::ConstantOp>(
              Loc, Builder.getF64FloatAttr(0.0));
        }
        Builder.create<scf::YieldOp>(Loc, ElseV);
      });

  if (CodegenFailed)
    return {};
  return IfOp.getResult(0);
}

Value ForExprAST::codegen() {
  // Emit the start value before putting the loop variable in scope.
  Value StartVal = Start->codegen();
  if (!StartVal)
    return {};

  Value Variable = CreateEntryBlockStorage();
  TheBuilder->create<memref::StoreOp>(getLocation(), StartVal, Variable,
                                      ValueRange{});

  auto OldValue = NamedValues.find(VarName);
  bool HadOldValue = OldValue != NamedValues.end();
  Value SavedValue = HadOldValue ? OldValue->second : Value();
  NamedValues[VarName] = Variable;
  bool CodegenFailed = false;

  // Test the condition before each iteration, then emit the body and step.
  TheBuilder->create<scf::WhileOp>(
      getLocation(), TypeRange{}, ValueRange{},
      [&](OpBuilder &Builder, Location Loc, ValueRange) {
        Value EndCond = End->codegen();
        if (!EndCond) {
          CodegenFailed = true;
          EndCond = Builder.create<arith::ConstantOp>(
              Loc, Builder.getF64FloatAttr(0.0));
        }

        Value Zero = Builder.create<arith::ConstantOp>(
            Loc, Builder.getF64FloatAttr(0.0));
        EndCond = Builder.create<arith::CmpFOp>(Loc, arith::CmpFPredicate::ONE,
                                                EndCond, Zero);
        Builder.create<scf::ConditionOp>(Loc, EndCond, ValueRange{});
      },
      [&](OpBuilder &Builder, Location Loc, ValueRange) {
        if (!Body->codegen())
          CodegenFailed = true;

        Value StepVal;
        if (Step)
          StepVal = Step->codegen();
        else
          StepVal = Builder.create<arith::ConstantOp>(
              Loc, Builder.getF64FloatAttr(1.0));
        if (!StepVal) {
          CodegenFailed = true;
          StepVal = Builder.create<arith::ConstantOp>(
              Loc, Builder.getF64FloatAttr(1.0));
        }

        // Reload after the body and step in case either mutated the variable.
        Value Current =
            Builder.create<memref::LoadOp>(Loc, Variable, ValueRange{});
        Value NextVar = Builder.create<arith::AddFOp>(Loc, Current, StepVal);
        Builder.create<memref::StoreOp>(Loc, NextVar, Variable, ValueRange{});
        Builder.create<scf::YieldOp>(Loc);
      });

  // Restore any variable shadowed by the loop induction variable.
  if (HadOldValue)
    NamedValues[VarName] = SavedValue;
  else
    NamedValues.erase(VarName);

  if (CodegenFailed)
    return {};

  // A for expression always returns 0.0.
  return TheBuilder->create<arith::ConstantOp>(
      getLocation(), TheBuilder->getF64FloatAttr(0.0));
}

Value VarExprAST::codegen() {
  std::vector<std::pair<std::string, std::optional<Value>>> OldBindings;

  auto RestoreBindings = [&]() {
    for (auto It = OldBindings.rbegin(); It != OldBindings.rend(); ++It) {
      if (It->second)
        NamedValues[It->first] = *It->second;
      else
        NamedValues.erase(It->first);
    }
  };

  for (auto &Variable : VarNames) {
    const std::string &Name = Variable.first;

    // Generate the initializer before introducing the new binding.
    Value InitialValue;
    if (Variable.second)
      InitialValue = Variable.second->codegen();
    else
      InitialValue = TheBuilder->create<arith::ConstantOp>(
          getLocation(), TheBuilder->getF64FloatAttr(0.0));
    if (!InitialValue) {
      RestoreBindings();
      return {};
    }

    Value Storage = CreateEntryBlockStorage();
    TheBuilder->create<memref::StoreOp>(getLocation(), InitialValue, Storage,
                                        ValueRange{});

    auto Old = NamedValues.find(Name);
    OldBindings.emplace_back(
        Name, Old == NamedValues.end() ? std::optional<Value>()
                                      : std::optional<Value>(Old->second));
    NamedValues[Name] = Storage;
  }

  Value BodyValue = Body->codegen();
  RestoreBindings();
  return BodyValue;
}

func::FuncOp PrototypeAST::codegen() {
  // Make the function type: double(double, double), etc.
  std::vector<Type> Doubles(Args.size(), TheBuilder->getF64Type());
  auto FunctionType =
      TheBuilder->getFunctionType(Doubles, {TheBuilder->getF64Type()});

  auto Function = func::FuncOp::create(getLocation(), Name, FunctionType);
  TheModule->push_back(Function);
  return Function;
}

func::FuncOp FunctionAST::codegen() {
  // Save the prototype so declarations can be emitted in later modules.
  auto &P = *Proto;
  FunctionProtos[Proto->getName()] = std::move(Proto);
  auto TheFunction = getFunction(P.getName());

  if (!TheFunction)
    return {};

  if (!TheFunction.isDeclaration()) {
    LogError("Function cannot be redefined.");
    return {};
  }

  // A definition is visible outside the module, even if an earlier extern
  // declaration created the function with private symbol visibility.
  TheFunction.setPublic();

  // If this is a binary operator, install its precedence.
  if (P.isBinaryOp())
    BinopPrecedence[P.getOperatorName()] = P.getBinaryPrecedence();

  // Create a new basic block to start insertion into.
  Block *EntryBlock = TheFunction.addEntryBlock();
  TheBuilder->setInsertionPointToStart(EntryBlock);

  // Give each function argument a mutable storage slot.
  NamedValues.clear();
  unsigned Index = 0;
  for (BlockArgument Argument : TheFunction.getArguments()) {
    Value Storage = CreateEntryBlockStorage();
    TheBuilder->create<memref::StoreOp>(getLocation(), Argument, Storage,
                                        ValueRange{});
    NamedValues[P.getArgs()[Index++]] = Storage;
  }

  if (Value RetVal = Body->codegen()) {
    // Finish off the function.
    TheBuilder->create<func::ReturnOp>(getLocation(), RetVal);

    // Validate the generated code, checking for consistency.
    if (succeeded(verify(TheFunction))) {
      // Run the optimizer on the module.
      if (failed(ThePM->run(*TheModule))) {
        LogError("Could not optimize function.");
        TheFunction.erase();
        if (P.isBinaryOp())
          BinopPrecedence.erase(P.getOperatorName());
        return {};
      }
      return TheFunction;
    }
  }

  // Error reading body, remove function.
  TheFunction.erase();
  if (P.isBinaryOp())
    BinopPrecedence.erase(P.getOperatorName());
  return {};
}

//===----------------------------------------------------------------------===//
// Top-Level parsing and JIT Driver
//===----------------------------------------------------------------------===//

static void InitializeModuleAndManagers() {
  // Destroy objects that refer to the old context before replacing it.
  ThePM.reset();
  TheBuilder.reset();
  TheModule = OwningOpRef<ModuleOp>();
  TheContext.reset();

  // Open a new context and module.
  TheContext = std::make_unique<MLIRContext>();
  TheContext->loadDialect<arith::ArithDialect, cf::ControlFlowDialect,
                          func::FuncDialect, memref::MemRefDialect,
                          scf::SCFDialect>();
  TheModule = ModuleOp::create(UnknownLoc::get(TheContext.get()));

  // Create a new builder for the module.
  TheBuilder = std::make_unique<OpBuilder>(TheContext.get());

  // Create a pass manager and add a couple of simple optimizations.
  ThePM = std::make_unique<PassManager>(TheContext.get());
  ThePM->addNestedPass<func::FuncOp>(createCanonicalizerPass());
  ThePM->addNestedPass<func::FuncOp>(createCSEPass());
}

static llvm::Expected<llvm::orc::ThreadSafeModule> lowerToLLVM() {
  // Lower the high-level MLIR operations to the LLVM dialect.
  PassManager LoweringPM(TheContext.get());
  LoweringPM.addPass(createSCFToControlFlowPass());
  LoweringPM.addPass(createMem2Reg());
  LoweringPM.addPass(createConvertFuncToLLVMPass());
  LoweringPM.addPass(createArithToLLVMConversionPass());
  LoweringPM.addPass(createFinalizeMemRefToLLVMConversionPass());
  LoweringPM.addPass(createConvertControlFlowToLLVMPass());

  // Clean up any temporary casts introduced by dialect conversion.
  LoweringPM.addPass(createReconcileUnrealizedCastsPass());
  if (failed(LoweringPM.run(*TheModule)))
    return llvm::make_error<llvm::StringError>(
        "could not lower module to the LLVM dialect",
        llvm::inconvertibleErrorCode());

  // Register the translations from MLIR's LLVM dialect to LLVM IR.
  registerBuiltinDialectTranslation(*TheContext);
  registerLLVMDialectTranslation(*TheContext);

  // Translate the lowered MLIR module into an LLVM IR module. The LLVM
  // context is kept with the module because the JIT may compile it later.
  auto LLVMContext = std::make_unique<llvm::LLVMContext>();
  auto LLVMModule = translateModuleToLLVMIR(*TheModule, *LLVMContext);
  if (!LLVMModule)
    return llvm::make_error<llvm::StringError>(
        "could not translate the LLVM dialect to LLVM IR",
        llvm::inconvertibleErrorCode());

  // Match the module's data layout to the target selected by the JIT.
  LLVMModule->setDataLayout(TheJIT->getDataLayout());

  if (DumpLLVMIR) {
    LLVMModule->print(llvm::errs(), nullptr);
    llvm::errs() << '\n';
  }

  // ThreadSafeModule transfers ownership of both objects to the ORC JIT.
  return llvm::orc::ThreadSafeModule(std::move(LLVMModule),
                                     std::move(LLVMContext));
}

static void HandleDefinition() {
  if (auto FnAST = ParseDefinition()) {
    if (auto FnIR = FnAST->codegen()) {
      if (DumpMLIR) {
        fprintf(stderr, "Read function definition:\n");
        FnIR.print(llvm::errs(), OpPrintingFlags().assumeVerified());
        fprintf(stderr, "\n");
      }

      ExitOnErr(TheJIT->addModule(ExitOnErr(lowerToLLVM())));
      InitializeModuleAndManagers();
    }
  } else {
    // Skip token for error recovery.
    getNextToken();
  }
}

static void HandleExtern() {
  if (auto ProtoAST = ParseExtern()) {
    if (auto FnIR = ProtoAST->codegen()) {
      FnIR.setPrivate();
      if (DumpMLIR) {
        fprintf(stderr, "Read extern:\n");
        FnIR.print(llvm::errs(), OpPrintingFlags().assumeVerified());
        fprintf(stderr, "\n");
      }
      FunctionProtos[ProtoAST->getName()] = std::move(ProtoAST);
    }
  } else {
    // Skip token for error recovery.
    getNextToken();
  }
}

static void HandleTopLevelExpression() {
  // Evaluate a top-level expression into an anonymous function.
  if (auto FnAST = ParseTopLevelExpr()) {
    if (auto FnIR = FnAST->codegen()) {
      if (DumpMLIR) {
        fprintf(stderr, "Read top-level expression:\n");
        FnIR.print(llvm::errs(), OpPrintingFlags().assumeVerified());
        fprintf(stderr, "\n");
      }

      auto RT = TheJIT->getMainJITDylib().createResourceTracker();
      ExitOnErr(TheJIT->addModule(ExitOnErr(lowerToLLVM()), RT));
      InitializeModuleAndManagers();

      auto ExprSymbol = ExitOnErr(TheJIT->lookup("__anon_expr"));
      double (*FP)() = ExprSymbol.getAddress().toPtr<double (*)()>();
      fprintf(stderr, "Evaluated to %f\n", FP());

      ExitOnErr(RT->remove());
    }
  } else {
    // Skip token for error recovery.
    getNextToken();
  }
}

//===----------------------------------------------------------------------===//
// "Library" functions that can be "extern'd" from user code.
//===----------------------------------------------------------------------===//

#ifdef _WIN32
#define DLLEXPORT __declspec(dllexport)
#else
#define DLLEXPORT
#endif

/// putchard - putchar that takes a double and returns 0.
extern "C" DLLEXPORT double putchard(double X) {
  fputc((char)X, stderr);
  return 0;
}

/// printd - printf that takes a double, prints it as "%f\n", and returns 0.
extern "C" DLLEXPORT double printd(double X) {
  fprintf(stderr, "%f\n", X);
  return 0;
}

/// top ::= definition | external | expression | ';'
static void MainLoop() {
  while (true) {
    switch (CurTok) {
    case tok_eof:
      return;
    case ';': // ignore top-level semicolons.
      fprintf(stderr, "ready> ");
      getNextToken();
      continue;
    case tok_def:
      HandleDefinition();
      break;
    case tok_extern:
      HandleExtern();
      break;
    default:
      HandleTopLevelExpression();
      break;
    }
    if (CurTok != tok_eof && CurTok != ';')
      fprintf(stderr, "ready> ");
  }
}

//===----------------------------------------------------------------------===//
// Main driver code.
//===----------------------------------------------------------------------===//

int main(int argc, char **argv) {
  llvm::cl::ParseCommandLineOptions(argc, argv, "Kaleidoscope JIT\n");

  llvm::InitializeNativeTarget();
  llvm::InitializeNativeTargetAsmPrinter();
  llvm::InitializeNativeTargetAsmParser();

  // Install standard binary operators.
  // 1 is lowest precedence.
  BinopPrecedence['='] = 2;
  BinopPrecedence['<'] = 10;
  BinopPrecedence['+'] = 20;
  BinopPrecedence['-'] = 20;
  BinopPrecedence['*'] = 40; // highest.

  // Prime the first token.
  fprintf(stderr, "ready> ");
  getNextToken();

  TheJIT = ExitOnErr(llvm::orc::KaleidoscopeJIT::Create());

  // Make the first module, which holds newly generated code.
  InitializeModuleAndManagers();

  // Run the main "interpreter loop" now.
  MainLoop();

  return 0;
}