9. Kaleidoscope: Adding Debug Information

9.1 Chapter 9 Introduction

Welcome to Chapter 9 of the "Implementing a language with MLIR" tutorial. In chapters 1 through 8, we've built a decent little programming language with functions and variables. What happens if something goes wrong though, how do you debug your program?

Source level debugging uses formatted data that helps a debugger translate from binary and the state of the machine back to the source that the programmer wrote. In LLVM we generally use a format called DWARF. DWARF is a compact encoding that represents types, source locations, and variable locations.

The short summary of this chapter is that we'll go through the various things you have to add to a programming language to support debug info, and how you translate that into DWARF.

Here's the sample program we'll be compiling:

# fib.ks
def fib(x)
  if x < 3 then
    1
  else
    fib(x-1)+fib(x-2);

fib(10)

9.2 Why is this a hard problem?

Debug information is a hard problem for a few different reasons - mostly centered around optimized code. First, optimization makes keeping source locations more difficult. MLIR operations carry source locations, which are preserved through lowering and ultimately become LLVM IR debug locations. Optimization passes should keep the source locations for newly created instructions, but merged instructions only get to keep a single location - this can cause jumping around when stepping through optimized programs. Secondly, optimization can move variables in ways that are either optimized out, shared in memory with other variables, or difficult to track.

Kaleidoscope enables optimization by default. When you want predictable source-level stepping, compile with -O0:

$ ./build/toy --emit-object -O0 fib.ks
Wrote fib.o

Use -O1, -O2, or -O3 to enable the canonicalization and CSE passes from Chapter 4 together with the corresponding LLVM code-generation optimization level. If you omit the option, Kaleidoscope uses -O2.

We accept the optimization level using LLVM's command-line library and default to -O2:

static llvm::cl::opt<char>
    OptLevel("O", llvm::cl::desc("Optimization level: -O0, -O1, -O2, or -O3"),
             llvm::cl::Prefix, llvm::cl::init('2'));

When initializing the MLIR pass manager, we add the optimization passes at every level above -O0:

// Create a pass manager and enable our simple optimizations above -O0.
ThePM = std::make_unique<PassManager>(TheContext.get());
if (OptLevel != '0') {
  ThePM->addNestedPass<func::FuncOp>(createCanonicalizerPass());
  ThePM->addNestedPass<func::FuncOp>(createCSEPass());
}

Finally, we pass the corresponding code-generation optimization level to LLVM's target machine:

llvm::TargetOptions Options;
llvm::CodeGenOptLevel CodeGenOpt;
switch (OptLevel) {
case '0': CodeGenOpt = llvm::CodeGenOptLevel::None; break;
case '1': CodeGenOpt = llvm::CodeGenOptLevel::Less; break;
case '2': CodeGenOpt = llvm::CodeGenOptLevel::Default; break;
case '3': CodeGenOpt = llvm::CodeGenOptLevel::Aggressive; break;
}
std::unique_ptr<llvm::TargetMachine> TargetMachine(
    Target->createTargetMachine(llvm::Triple(TargetTriple), "generic", "",
                                Options, llvm::Reloc::PIC_, std::nullopt,
                                CodeGenOpt));

9.3 Compile Unit

The top level container for a section of code in DWARF is a compile unit. This contains the type and function data for an individual translation unit (read: one file of source code). The debug information for the functions in fib.ks, for example, belongs to a single compile unit.

9.4 Source Locations

The most important thing for debug information is accurate source location - this makes it possible to map your source code back. We have a problem though, Kaleidoscope really doesn't have any source location information in the lexer or parser so we'll need to add it.

struct SourceLocation {
  int Line;
  int Col;
};
static SourceLocation CurLoc;
static SourceLocation LexLoc = {1, 0};

static int advance() {
  int LastChar = getchar();
  if (LastChar == '\n' || LastChar == '\r') {
    ++LexLoc.Line;
    LexLoc.Col = 0;
  } else {
    ++LexLoc.Col;
  }
  return LastChar;
}

In this set of code we've added some functionality on how to keep track of the line and column of the "source file". As we lex every token we set our current "lexical location" to the assorted line and column for the beginning of the token. We do this by overriding all of the previous calls to getchar() with our new advance() that keeps track of the information and then we have added to all of our AST classes a source location:

class ExprAST {
  SourceLocation Loc;

public:
  ExprAST(SourceLocation Loc = CurLoc) : Loc(Loc) {}
  virtual ~ExprAST() = default;

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

We turn the current AST location into a FileLineColLoc with this helper:

static SourceLocation CodegenLoc = {1, 1};

static Location getLocation() {
  llvm::StringRef Filename = "<stdin>";
  if (!InputFilename.empty())
    Filename = InputFilename.getValue();
  return FileLineColLoc::get(TheContext.get(), Filename, CodegenLoc.Line,
                             CodegenLoc.Col);
}

A small guard selects an AST node's location while its code is generated and restores the previous location afterward:

class LocationGuard {
  SourceLocation Previous;

public:
  explicit LocationGuard(SourceLocation Loc) : Previous(CodegenLoc) {
    CodegenLoc = Loc;
  }
  ~LocationGuard() { CodegenLoc = Previous; }
};

Each code-generation method creates one of these guards before it creates any MLIR operations. For example:

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

The location is now part of the MLIR operation. Conversion passes carry it through the lowering pipeline, and the final LLVM translation turns it into a DILocation attached to the generated instruction.

9.5 Variables and Our First Custom Pass

Now that we have functions, we need to be able to print out the variables we have in scope. Let's get our function arguments set up so we can get decent backtraces and see how our functions are being called. MLIR represents values in SSA form and does not preserve our source-level variable names. A source location tells the debugger where an operation came from, but not that a value represents a variable named x. To make function arguments visible by name in the debugger, we need to attach that information explicitly.

FunctionParameters[P.getName()] = P.getArgs();

We could walk the lowered module and add the remaining debug attributes here, but that work is a natural fit for an MLIR pass. A pass is a structured way to inspect or transform some IR, and frontends can add their own passes alongside the ones supplied by MLIR.

Our KaleidoscopeDebugInfoPass runs after lowering to the LLVM dialect. It creates the compile unit, describes each function and its parameters, and connects those parameters to their lowered stack storage. DWARF has no language identifier for Kaleidoscope, so this tutorial records C as a practical stand-in. The implementation lives in KaleidoscopeDebugInfo.cpp for readers interested in the debug metadata itself. The general machinery for writing a pass is described in MLIR's Pass Infrastructure documentation.

The pass is exposed through a small creation function:

std::unique_ptr<mlir::Pass> createKaleidoscopeDebugInfoPass(
    llvm::StringRef inputFilename, char optLevel,
    const std::map<std::string, std::vector<std::string>> &functionParameters);

We add it to a pass manager just like any built-in MLIR pass:

PassManager DebugPM(TheContext.get());
DebugPM.addPass(createKaleidoscopeDebugInfoPass(
    InputFilename.getValue(), OptLevel, FunctionParameters));

We follow it with MLIR's ensure debug info scope on LLVM functions pass. Our pass supplies the language-specific compile unit, functions, and variables; MLIR's pass preserves that information and fills in the remaining scopes on the lowered operations:

LLVM::DIScopeForLLVMFuncOpPassOptions DebugOptions;
DebugOptions.emissionKind = LLVM::DIEmissionKind::Full;
DebugPM.addPass(
    LLVM::createDIScopeForLLVMFuncOpPass(std::move(DebugOptions)));
if (failed(DebugPM.run(*TheModule)))
  return llvm::make_error<llvm::StringError>(
      "could not add LLVM debug scopes",
      llvm::inconvertibleErrorCode());

With this we have enough information to set breakpoints in functions, print their arguments, and inspect the call stack. Translation from the LLVM dialect then produces the final LLVM debug metadata.

9.6 Inspecting the Debug Information

We can now compile fib.ks to an object file:

$ ./build/toy --emit-object -O0 fib.ks
Wrote fib.o

The llvm-dwarfdump tool lets us inspect the DWARF information in that object. On my Apple silicon Mac, the relevant part of the output looks like this:

$ llvm-dwarfdump --debug-info fib.o
fib.o: file format Mach-O arm64
...
0x00000026:   DW_TAG_subprogram
                DW_AT_low_pc              (0x0000000000000000)
                DW_AT_high_pc             (0x0000000000000090)
                DW_AT_APPLE_omit_frame_ptr (true)
                DW_AT_frame_base          (DW_OP_reg31 WSP)
                DW_AT_linkage_name        ("fib")
                DW_AT_name                ("fib")
                DW_AT_decl_file           ("fib.ks")
                DW_AT_decl_line           (2)
                DW_AT_external            (true)

0x0000003f:     DW_TAG_formal_parameter
                  DW_AT_location  (DW_OP_fbreg +24)
                  DW_AT_name      ("x")
                  DW_AT_decl_file ("fib.ks")
                  DW_AT_decl_line (2)
                  DW_AT_type      (0x00000067 "double")

0x00000067:   DW_TAG_base_type
                DW_AT_name      ("double")
                DW_AT_encoding  (DW_ATE_float)
                DW_AT_byte_size (0x08)
...

Offsets, addresses, and some target-specific attributes will differ between systems. The important parts are the fib subprogram, its source file and line, and the parameter named x with type double.

9.7 Download Source Code

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

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

9.8 Full Code Listing

Here is the complete code listing for our running example, enhanced with debug information. Here is the CMake configuration:

cmake_minimum_required(VERSION 3.20)

project(kaleidoscope-chapter-09 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
  KaleidoscopeDebugInfo.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
  CodeGen
  OrcJIT
  Support
  Target
  ${LLVM_TARGETS_TO_BUILD}
)

target_link_libraries(toy PRIVATE
  MLIRArithToLLVM
  MLIRArithDialect
  MLIRBuiltinToLLVMIRTranslation
  MLIRControlFlowDialect
  MLIRControlFlowToLLVM
  MLIRFuncToLLVM
  MLIRFuncDialect
  MLIRLLVMDialect
  MLIRLLVMIRTransforms
  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 "KaleidoscopeDebugInfo.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/LLVMIR/LLVMDialect.h"
#include "mlir/Dialect/LLVMIR/Transforms/Passes.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/ADT/SmallString.h"
#include "llvm/IR/LLVMContext.h"
#include "llvm/IR/LegacyPassManager.h"
#include "llvm/IR/Module.h"
#include "llvm/MC/TargetRegistry.h"
#include "llvm/Support/CommandLine.h"
#include "llvm/Support/Error.h"
#include "llvm/Support/FileSystem.h"
#include "llvm/Support/Path.h"
#include "llvm/Support/TargetSelect.h"
#include "llvm/Support/raw_ostream.h"
#include "llvm/Target/TargetMachine.h"
#include "llvm/Target/TargetOptions.h"
#include "llvm/TargetParser/Host.h"
#include <cassert>
#include <cctype>
#include <cstdio>
#include <cstdlib>
#include <map>
#include <memory>
#include <optional>
#include <string>
#include <system_error>
#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
};

struct SourceLocation {
  int Line;
  int Col;
};

static SourceLocation CurLoc;
static SourceLocation LexLoc = {1, 0};

static int advance() {
  int LastChar = getchar();
  if (LastChar == '\n' || LastChar == '\r') {
    ++LexLoc.Line;
    LexLoc.Col = 0;
  } else {
    ++LexLoc.Col;
  }
  return LastChar;
}

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 = advance();

  CurLoc = LexLoc;

  if (isalpha(LastChar)) { // identifier: [a-zA-Z][a-zA-Z0-9]*
    IdentifierStr = LastChar;
    while (isalnum((LastChar = advance())))
      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 = advance();
    } while (isdigit(LastChar) || LastChar == '.');

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

  if (LastChar == '#') {
    // Comment until end of line.
    do
      LastChar = advance();
    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 = advance();
  return ThisChar;
}

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

namespace {

/// ExprAST - Base class for all expression nodes.
class ExprAST {
  SourceLocation Loc;

public:
  ExprAST(SourceLocation Loc = CurLoc) : Loc(Loc) {}
  virtual ~ExprAST() = default;

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

/// 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(SourceLocation Loc, const std::string &Name)
      : ExprAST(Loc), 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(SourceLocation Loc, char Opcode,
               std::unique_ptr<ExprAST> Operand)
      : ExprAST(Loc), 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(SourceLocation Loc, char Op, std::unique_ptr<ExprAST> LHS,
                std::unique_ptr<ExprAST> RHS)
      : ExprAST(Loc), 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(SourceLocation Loc, const std::string &Callee,
              std::vector<std::unique_ptr<ExprAST>> Args)
      : ExprAST(Loc), 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(SourceLocation Loc, std::unique_ptr<ExprAST> Cond,
            std::unique_ptr<ExprAST> Then, std::unique_ptr<ExprAST> Else)
      : ExprAST(Loc), 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(SourceLocation Loc, const std::string &VarName,
             std::unique_ptr<ExprAST> Start,
             std::unique_ptr<ExprAST> End, std::unique_ptr<ExprAST> Step,
             std::unique_ptr<ExprAST> Body)
      : ExprAST(Loc), 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(SourceLocation Loc,
      std::vector<std::pair<std::string, std::unique_ptr<ExprAST>>> VarNames,
      std::unique_ptr<ExprAST> Body)
      : ExprAST(Loc), 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.
  SourceLocation Loc;

public:
  PrototypeAST(SourceLocation Loc, 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), Loc(Loc) {}

  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; }
  SourceLocation getSourceLocation() const { return Loc; }
};

/// 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;
  SourceLocation IdLoc = CurLoc;

  getNextToken(); // eat identifier.

  if (CurTok != '(') // Simple variable ref.
    return std::make_unique<VariableExprAST>(IdLoc, 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>(IdLoc, IdName, std::move(Args));
}

/// ifexpr ::= 'if' expression 'then' expression 'else' expression
static std::unique_ptr<ExprAST> ParseIfExpr() {
  SourceLocation IfLoc = CurLoc;
  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>(IfLoc, std::move(Cond), std::move(Then),
                                     std::move(Else));
}

/// forexpr ::= 'for' identifier '=' expr ',' expr (',' expr)? 'in' expression
static std::unique_ptr<ExprAST> ParseForExpr() {
  SourceLocation ForLoc = CurLoc;
  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>(ForLoc, 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() {
  SourceLocation VarLoc = CurLoc;
  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>(VarLoc, 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;
  SourceLocation UnaryLoc = CurLoc;
  getNextToken();
  if (auto Operand = ParseUnary())
    return std::make_unique<UnaryExprAST>(UnaryLoc, 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;
    SourceLocation BinLoc = CurLoc;
    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>(BinLoc, 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() {
  SourceLocation FnLoc = CurLoc;
  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>(FnLoc, 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() {
  SourceLocation FnLoc = CurLoc;
  if (auto E = ParseExpression()) {
    // Make an anonymous proto.
    auto Proto = std::make_unique<PrototypeAST>(FnLoc, "__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 std::map<std::string, std::vector<std::string>> FunctionParameters;
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>
    EmitObject("emit-object",
               llvm::cl::desc(
                   "Compile an input file to an object instead of using the JIT"),
               llvm::cl::init(false));
static llvm::cl::opt<std::string>
    InputFilename(llvm::cl::Positional, llvm::cl::desc("<input file>"),
                  llvm::cl::init(""));
static llvm::cl::opt<std::string>
    TargetTripleOption("target",
                       llvm::cl::desc("Target triple for object emission"),
                       llvm::cl::value_desc("triple"), llvm::cl::init(""));
static llvm::cl::opt<std::string>
    OutputFilename("o", llvm::cl::desc("Output filename"),
                   llvm::cl::value_desc("filename"), llvm::cl::init(""));
static llvm::cl::opt<char>
    OptLevel("O", llvm::cl::desc("Optimization level: -O0, -O1, -O2, or -O3"),
             llvm::cl::Prefix, llvm::cl::init('2'));
static llvm::cl::opt<bool>
    DumpLLVMIR("dump-llvm-ir",
             llvm::cl::desc("Print LLVM IR before emitting the object file"),
             llvm::cl::init(false));

static SourceLocation CodegenLoc = {1, 1};

static Location getLocation() {
  llvm::StringRef Filename = "<stdin>";
  if (!InputFilename.empty())
    Filename = InputFilename.getValue();
  return FileLineColLoc::get(TheContext.get(), Filename, CodegenLoc.Line,
                             CodegenLoc.Col);
}

class LocationGuard {
  SourceLocation Previous;

public:
  explicit LocationGuard(SourceLocation Loc) : Previous(CodegenLoc) {
    CodegenLoc = Loc;
  }
  ~LocationGuard() { CodegenLoc = Previous; }
};

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() {
  LocationGuard Guard(getSourceLocation());
  return TheBuilder->create<arith::ConstantOp>(
      getLocation(), TheBuilder->getF64FloatAttr(Val));
}

Value VariableExprAST::codegen() {
  LocationGuard Guard(getSourceLocation());
  // 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() {
  LocationGuard Guard(getSourceLocation());
  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() {
  LocationGuard Guard(getSourceLocation());
  // 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() {
  LocationGuard Guard(getSourceLocation());
  // 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() {
  LocationGuard Guard(getSourceLocation());
  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() {
  LocationGuard Guard(getSourceLocation());
  // 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() {
  LocationGuard Guard(getSourceLocation());
  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() {
  LocationGuard Guard(getSourceLocation());
  // 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;
  LocationGuard Guard(P.getSourceLocation());
  FunctionParameters[P.getName()] = P.getArgs();
  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 enable our simple optimizations above -O0.
  ThePM = std::make_unique<PassManager>(TheContext.get());
  if (OptLevel != '0') {
    ThePM->addNestedPass<func::FuncOp>(createCanonicalizerPass());
    ThePM->addNestedPass<func::FuncOp>(createCSEPass());
  }
}

struct LoweredModule {
  std::unique_ptr<llvm::LLVMContext> Context;
  std::unique_ptr<llvm::Module> Module;
};

static llvm::Expected<LoweredModule>
lowerToLLVM(const llvm::DataLayout &DataLayout) {
  // Lower the high-level MLIR operations to the LLVM dialect.
  PassManager LoweringPM(TheContext.get());
  LoweringPM.addPass(createSCFToControlFlowPass());
  if (OptLevel != '0')
    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());

  // Add our language-specific compile-unit, function, and parameter debug info.
  PassManager DebugPM(TheContext.get());
  DebugPM.addPass(createKaleidoscopeDebugInfoPass(
      InputFilename.getValue(), OptLevel, FunctionParameters));

  // Fill in the remaining debug scopes on the lowered LLVM operations.
  LLVM::DIScopeForLLVMFuncOpPassOptions DebugOptions;
  DebugOptions.emissionKind = LLVM::DIEmissionKind::Full;
  DebugPM.addPass(
      LLVM::createDIScopeForLLVMFuncOpPass(std::move(DebugOptions)));
  if (failed(DebugPM.run(*TheModule)))
    return llvm::make_error<llvm::StringError>(
        "could not add LLVM debug scopes",
        llvm::inconvertibleErrorCode());

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

  // Translate the lowered MLIR module into the LLVM IR module consumed by the
  // target's object-file emitter.
  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());

  LLVMModule->setDataLayout(DataLayout);

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

  return LoweredModule{std::move(LLVMContext), std::move(LLVMModule)};
}

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");
      }

      if (!EmitObject) {
        auto Lowered = ExitOnErr(lowerToLLVM(TheJIT->getDataLayout()));
        ExitOnErr(TheJIT->addModule(llvm::orc::ThreadSafeModule(
            std::move(Lowered.Module), std::move(Lowered.Context))));
        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 with the JIT, or retain it when compiling.
  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");
      }

      if (!EmitObject) {
        auto RT = TheJIT->getMainJITDylib().createResourceTracker();
        auto Lowered = ExitOnErr(lowerToLLVM(TheJIT->getDataLayout()));
        ExitOnErr(TheJIT->addModule(
            llvm::orc::ThreadSafeModule(std::move(Lowered.Module),
                                        std::move(Lowered.Context)),
            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.
      if (!EmitObject)
        fprintf(stderr, "ready> ");
      getNextToken();
      continue;
    case tok_def:
      HandleDefinition();
      break;
    case tok_extern:
      HandleExtern();
      break;
    default:
      HandleTopLevelExpression();
      break;
    }
    if (!EmitObject && CurTok != tok_eof && CurTok != ';')
      fprintf(stderr, "ready> ");
  }
}

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

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

  if (OptLevel < '0' || OptLevel > '3') {
    llvm::errs() << "Error: optimization level must be -O0, -O1, -O2, or -O3\n";
    return 1;
  }

  if (EmitObject && InputFilename.empty()) {
    llvm::errs() << "Error: --emit-object requires an input file\n";
    return 1;
  }
  if (!EmitObject && !InputFilename.empty()) {
    llvm::errs() << "Error: an input file requires --emit-object\n";
    return 1;
  }
  if (!EmitObject && !TargetTripleOption.empty()) {
    llvm::errs() << "Error: --target requires --emit-object\n";
    return 1;
  }
  if (!EmitObject && !OutputFilename.empty()) {
    llvm::errs() << "Error: -o requires --emit-object\n";
    return 1;
  }
  if (EmitObject && !std::freopen(InputFilename.c_str(), "r", stdin)) {
    std::perror(("Error opening " + InputFilename).c_str());
    return 1;
  }

  llvm::InitializeAllTargetInfos();
  llvm::InitializeAllTargets();
  llvm::InitializeAllTargetMCs();
  llvm::InitializeAllAsmParsers();
  llvm::InitializeAllAsmPrinters();

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

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

  if (!EmitObject)
    TheJIT = ExitOnErr(llvm::orc::KaleidoscopeJIT::Create());

  // JIT mode replaces this module as definitions are submitted. Object mode
  // retains it until the entire input has been parsed.
  InitializeModuleAndManagers();

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

  if (!EmitObject)
    return 0;

  // Select the host target and configure its object-file emitter.
  std::string TargetTriple =
      TargetTripleOption.empty()
          ? llvm::sys::getDefaultTargetTriple()
          : llvm::Triple::normalize(TargetTripleOption);
  std::string Error;
  const llvm::Target *Target =
      llvm::TargetRegistry::lookupTarget(TargetTriple, Error);
  if (!Target) {
    llvm::errs() << Error << '\n';
    return 1;
  }

  llvm::TargetOptions Options;
  llvm::CodeGenOptLevel CodeGenOpt;
  switch (OptLevel) {
  case '0': CodeGenOpt = llvm::CodeGenOptLevel::None; break;
  case '1': CodeGenOpt = llvm::CodeGenOptLevel::Less; break;
  case '2': CodeGenOpt = llvm::CodeGenOptLevel::Default; break;
  case '3': CodeGenOpt = llvm::CodeGenOptLevel::Aggressive; break;
  }
  std::unique_ptr<llvm::TargetMachine> TargetMachine(
      Target->createTargetMachine(llvm::Triple(TargetTriple), "generic", "",
                                  Options, llvm::Reloc::PIC_, std::nullopt,
                                  CodeGenOpt));
  if (!TargetMachine) {
    llvm::errs() << "Could not create the target machine\n";
    return 1;
  }

  // Lower the complete MLIR module once, then attach the target information
  // required to produce a native object file.
  auto Lowered = ExitOnErr(lowerToLLVM(TargetMachine->createDataLayout()));
  Lowered.Module->setTargetTriple(llvm::Triple(TargetTriple));

  llvm::SmallString<256> Filename;
  if (OutputFilename.empty()) {
    Filename = InputFilename;
    llvm::sys::path::replace_extension(Filename, "o");
  } else {
    Filename = OutputFilename;
  }
  std::error_code EC;
  llvm::raw_fd_ostream Dest(Filename, EC, llvm::sys::fs::OF_None);
  if (EC) {
    llvm::errs() << "Could not open " << Filename << ": " << EC.message()
                 << '\n';
    return 1;
  }

  llvm::legacy::PassManager EmitPM;
  if (TargetMachine->addPassesToEmitFile(
          EmitPM, Dest, nullptr, llvm::CodeGenFileType::ObjectFile)) {
    llvm::errs() << "Target machine cannot emit an object file\n";
    return 1;
  }

  EmitPM.run(*Lowered.Module);
  Dest.flush();
  llvm::outs() << "Wrote " << Filename << '\n';

  return 0;
}

Here is the interface for our debug-information pass:

#ifndef KALEIDOSCOPE_DEBUG_INFO_H
#define KALEIDOSCOPE_DEBUG_INFO_H

#include "mlir/Pass/Pass.h"
#include "llvm/ADT/StringRef.h"
#include <map>
#include <memory>
#include <string>
#include <vector>

/// Create the module pass that adds the debug metadata needed before the LLVM
/// dialect is translated to LLVM IR.
///
/// `inputFilename` identifies the source file (or is empty for stdin), and
/// `optLevel` determines whether the compile unit is marked as optimized.
/// `functionParameters` preserves source parameter names, which are no longer
/// present in the lowered LLVM function arguments themselves.
std::unique_ptr<mlir::Pass> createKaleidoscopeDebugInfoPass(
    llvm::StringRef inputFilename, char optLevel,
    const std::map<std::string, std::vector<std::string>> &functionParameters);

#endif

And here is its implementation:

#include "KaleidoscopeDebugInfo.h"
#include "mlir/Dialect/LLVMIR/LLVMDialect.h"
#include "mlir/IR/Builders.h"
#include "mlir/IR/BuiltinOps.h"
#include "llvm/BinaryFormat/Dwarf.h"
#include "llvm/Support/Path.h"

using namespace mlir;

namespace {

using FunctionParameterMap = std::map<std::string, std::vector<std::string>>;

/// Walk backwards from a value produced during lowering until we reach the
/// llvm.alloca that owns its stack storage.
///
/// A debug declaration describes the storage for a variable, not a temporary
/// value derived from it. Lowering may insert address calculations, bitcasts,
/// or other operations between the source variable and its alloca, so we
/// follow the operand chain until we find the alloca itself. Returns a null
/// Value if no alloca is found.
static Value findUnderlyingAlloca(Value value) {
  Operation *definingOp = value.getDefiningOp();
  if (!definingOp)
    return {};
  if (isa<LLVM::AllocaOp>(definingOp))
    return value;
  for (Value operand : definingOp->getOperands())
    if (Value alloca = findUnderlyingAlloca(operand))
      return alloca;
  return {};
}

class KaleidoscopeDebugInfoPass
    : public PassWrapper<KaleidoscopeDebugInfoPass, OperationPass<ModuleOp>> {
public:
  MLIR_DEFINE_EXPLICIT_INTERNAL_INLINE_TYPE_ID(KaleidoscopeDebugInfoPass)

  KaleidoscopeDebugInfoPass(StringRef inputFilename, char optLevel,
                            const FunctionParameterMap &functionParameters)
      : inputFilename(inputFilename.str()), optLevel(optLevel),
        functionParameters(functionParameters) {}

  StringRef getArgument() const final { return "kaleidoscope-debug-info"; }
  StringRef getDescription() const final {
    return "Add Kaleidoscope compile-unit, function, and parameter debug info";
  }

  void runOnOperation() override {
    ModuleOp module = getOperation();
    MLIRContext *context = module.getContext();

    // A DWARF file entry is a pair of (directory, filename), not a single
    // path. Input read interactively from the REPL has no real path at all,
    // so we invent a synthetic <stdin> file to describe it.
    StringRef inputPath = inputFilename;
    auto file = inputPath.empty()
                    ? LLVM::DIFileAttr::get(context, "<stdin>", "")
                    : LLVM::DIFileAttr::get(
                          context, llvm::sys::path::filename(inputPath),
                          llvm::sys::path::parent_path(inputPath));

    // The compile unit is the top-level debug record for this translation
    // unit. Wrapping it in a DistinctAttr guarantees it will not be merged
    // with another compile unit that happens to have identical fields, which
    // would be a DWARF-validity bug.
    auto compileUnit = LLVM::DICompileUnitAttr::get(
        DistinctAttr::create(UnitAttr::get(context)), llvm::dwarf::DW_LANG_C,
        file, StringAttr::get(context, "Kaleidoscope"),
        /*isOptimized=*/optLevel != '0', LLVM::DIEmissionKind::Full);

    // MLIR carries debug metadata on locations rather than on a separate
    // side table. Fusing the compile unit into the module's location is how
    // we attach it so that LLVM IR translation can recover it later.
    module->setLoc(FusedLoc::get(context, {module.getLoc()}, compileUnit));

    // Kaleidoscope has exactly one source-level type today. Every parameter
    // we describe below uses this same DWARF base type.
    auto doubleType =
        LLVM::DIBasicTypeAttr::get(context, llvm::dwarf::DW_TAG_base_type,
                                   "double", 64, llvm::dwarf::DW_ATE_float);
    auto emptyExpression = LLVM::DIExpressionAttr::get(context);

    for (LLVM::LLVMFuncOp function : module.getOps<LLVM::LLVMFuncOp>()) {
      // Choose the DIFile and source line for this function, preferring the
      // information recorded on the function's source location. When the
      // parser supplied no concrete location, fall back to the module's input
      // file and line 1.
      Location originalLoc = function.getLoc();
      LLVM::DIFileAttr functionFile = file;
      int64_t line = 1;
      if (auto fileLoc = originalLoc->findInstanceOf<FileLineColLoc>()) {
        StringRef functionPath = fileLoc.getFilename().getValue();
        functionFile = LLVM::DIFileAttr::get(
            context, llvm::sys::path::filename(functionPath),
            llvm::sys::path::parent_path(functionPath));
        line = fileLoc.getLine();
      }

      // A definition and an external declaration need different DISubprogram
      // configurations. A definition belongs to this compile unit, carries a
      // distinct identity so it will not be merged with an identical-looking
      // subprogram, and is marked as a Definition. A declaration has no body
      // emitted by this compile unit, so it gets neither attachment.
      DistinctAttr id;
      LLVM::DICompileUnitAttr functionCompileUnit = compileUnit;
      auto flags = static_cast<LLVM::DISubprogramFlags>(0);
      if (optLevel != '0')
        flags = flags | LLVM::DISubprogramFlags::Optimized;
      if (function.isExternal()) {
        functionCompileUnit = {};
      } else {
        id = DistinctAttr::create(UnitAttr::get(context));
        flags = flags | LLVM::DISubprogramFlags::Definition;
      }

      // Kaleidoscope has no source-level function signatures yet, so we
      // cannot describe parameter or return types accurately here. An empty
      // DISubroutineType is the minimum required to attach a DISubprogram;
      // the parameters themselves are described later by DILocalVariable
      // entries on their stack slots.
      auto functionType = LLVM::DISubroutineTypeAttr::get(
          context, llvm::dwarf::DW_CC_normal, {});
      auto name = function.getNameAttr();
      auto scope = LLVM::DISubprogramAttr::get(
          context, id, functionCompileUnit, functionFile, name, name,
          functionFile, line, line, flags, functionType,
          /*retainedNodes=*/{}, /*annotations=*/{});

      // Fuse the subprogram scope into the function's existing location.
      // Later lowering and LLVM IR translation will use this scope for
      // instructions emitted inside the function body.
      function->setLoc(FusedLoc::get(context, {originalLoc}, scope));

      // Declarations have no body, and a definition may be absent from the
      // parser-side parameter table when it was generated by the compiler
      // rather than written by the user. In the latter case, the function may
      // still have parameters, but this pass does not know their source names.
      auto names = functionParameters.find(function.getName().str());
      if (function.isExternal() || names == functionParameters.end())
        continue;

      Block &entryBlock = function.getBody().front();
      for (auto [argumentNumber, parameterName] :
           llvm::enumerate(names->second)) {
        if (argumentNumber >= entryBlock.getNumArguments())
          break;

        BlockArgument argument = entryBlock.getArgument(argumentNumber);

        // A mutable Kaleidoscope parameter is copied into a stack slot at
        // the top of the entry block. Locate the store of the incoming
        // argument so we can point the debug declaration at the stack slot
        // rather than at the argument itself.
        LLVM::StoreOp argumentStore;
        for (LLVM::StoreOp store : entryBlock.getOps<LLVM::StoreOp>()) {
          if (store.getValue() == argument) {
            argumentStore = store;
            break;
          }
        }
        if (!argumentStore)
          continue;

        // DWARF argument numbers are one-based. The source name comes from
        // the parser-side table because LLVM block arguments do not retain
        // it. The scope, file, and line are inherited from the enclosing
        // function.
        auto variable = LLVM::DILocalVariableAttr::get(
            scope, parameterName, scope.getFile(), scope.getLine(),
            argumentNumber + 1, /*alignInBits=*/0, doubleType,
            LLVM::DIFlags::Zero);

        // The store's address operand may be an intermediate value rather
        // than the alloca itself. Walk back to the alloca if one exists;
        // otherwise, fall back to the store's own address.
        Value variableAddress = findUnderlyingAlloca(argumentStore.getAddr());
        if (!variableAddress)
          variableAddress = argumentStore.getAddr();

        // llvm.dbg.declare attaches the debug variable to the storage
        // address. It emits no executable code — it is a hint for the
        // debugger, translated to a #dbg_declare record in LLVM IR.
        OpBuilder builder(argumentStore);
        builder.setInsertionPointAfter(argumentStore);
        builder.create<LLVM::DbgDeclareOp>(
            argumentStore.getLoc(), variableAddress, variable, emptyExpression);
      }
    }
  }

private:
  // The pass owns copies of these because the caller's strings and maps are
  // not guaranteed to outlive execution of the pass manager.
  std::string inputFilename;
  char optLevel;
  FunctionParameterMap functionParameters;
};

} // namespace

// Construction is hidden behind a factory so toy.cpp does not need to know
// the concrete pass implementation type.
std::unique_ptr<Pass> createKaleidoscopeDebugInfoPass(
    StringRef inputFilename, char optLevel,
    const std::map<std::string, std::vector<std::string>> &functionParameters) {
  return std::make_unique<KaleidoscopeDebugInfoPass>(inputFilename, optLevel,
                                                     functionParameters);
}