1 #include "llvm/ADT/STLExtras.h" 2 #include "llvm/IR/IRBuilder.h" 3 #include "llvm/IR/LLVMContext.h" 4 #include "llvm/IR/Module.h" 5 #include "llvm/IR/Verifier.h" 6 #include <cctype> 7 #include <cstdio> 8 #include <map> 9 #include <string> 10 #include <vector> 11 12 using namespace llvm; 13 14 //===----------------------------------------------------------------------===// 15 // Lexer 16 //===----------------------------------------------------------------------===// 17 18 // The lexer returns tokens [0-255] if it is an unknown character, otherwise one 19 // of these for known things. 20 enum Token { 21 tok_eof = -1, 22 23 // commands 24 tok_def = -2, 25 tok_extern = -3, 26 27 // primary 28 tok_identifier = -4, 29 tok_number = -5 30 }; 31 32 static std::string IdentifierStr; // Filled in if tok_identifier 33 static double NumVal; // Filled in if tok_number 34 35 /// gettok - Return the next token from standard input. 36 static int gettok() { 37 static int LastChar = ' '; 38 39 // Skip any whitespace. 40 while (isspace(LastChar)) 41 LastChar = getchar(); 42 43 if (isalpha(LastChar)) { // identifier: [a-zA-Z][a-zA-Z0-9]* 44 IdentifierStr = LastChar; 45 while (isalnum((LastChar = getchar()))) 46 IdentifierStr += LastChar; 47 48 if (IdentifierStr == "def") 49 return tok_def; 50 if (IdentifierStr == "extern") 51 return tok_extern; 52 return tok_identifier; 53 } 54 55 if (isdigit(LastChar) || LastChar == '.') { // Number: [0-9.]+ 56 std::string NumStr; 57 do { 58 NumStr += LastChar; 59 LastChar = getchar(); 60 } while (isdigit(LastChar) || LastChar == '.'); 61 62 NumVal = strtod(NumStr.c_str(), nullptr); 63 return tok_number; 64 } 65 66 if (LastChar == '#') { 67 // Comment until end of line. 68 do 69 LastChar = getchar(); 70 while (LastChar != EOF && LastChar != '\n' && LastChar != '\r'); 71 72 if (LastChar != EOF) 73 return gettok(); 74 } 75 76 // Check for end of file. Don't eat the EOF. 77 if (LastChar == EOF) 78 return tok_eof; 79 80 // Otherwise, just return the character as its ascii value. 81 int ThisChar = LastChar; 82 LastChar = getchar(); 83 return ThisChar; 84 } 85 86 //===----------------------------------------------------------------------===// 87 // Abstract Syntax Tree (aka Parse Tree) 88 //===----------------------------------------------------------------------===// 89 namespace { 90 /// ExprAST - Base class for all expression nodes. 91 class ExprAST { 92 public: 93 virtual ~ExprAST() {} 94 virtual Value *codegen() = 0; 95 }; 96 97 /// NumberExprAST - Expression class for numeric literals like "1.0". 98 class NumberExprAST : public ExprAST { 99 double Val; 100 101 public: 102 NumberExprAST(double Val) : Val(Val) {} 103 Value *codegen() override; 104 }; 105 106 /// VariableExprAST - Expression class for referencing a variable, like "a". 107 class VariableExprAST : public ExprAST { 108 std::string Name; 109 110 public: 111 VariableExprAST(const std::string &Name) : Name(Name) {} 112 Value *codegen() override; 113 }; 114 115 /// BinaryExprAST - Expression class for a binary operator. 116 class BinaryExprAST : public ExprAST { 117 char Op; 118 std::unique_ptr<ExprAST> LHS, RHS; 119 120 public: 121 BinaryExprAST(char Op, std::unique_ptr<ExprAST> LHS, 122 std::unique_ptr<ExprAST> RHS) 123 : Op(Op), LHS(std::move(LHS)), RHS(std::move(RHS)) {} 124 Value *codegen() override; 125 }; 126 127 /// CallExprAST - Expression class for function calls. 128 class CallExprAST : public ExprAST { 129 std::string Callee; 130 std::vector<std::unique_ptr<ExprAST>> Args; 131 132 public: 133 CallExprAST(const std::string &Callee, 134 std::vector<std::unique_ptr<ExprAST>> Args) 135 : Callee(Callee), Args(std::move(Args)) {} 136 Value *codegen() override; 137 }; 138 139 /// PrototypeAST - This class represents the "prototype" for a function, 140 /// which captures its name, and its argument names (thus implicitly the number 141 /// of arguments the function takes). 142 class PrototypeAST { 143 std::string Name; 144 std::vector<std::string> Args; 145 146 public: 147 PrototypeAST(const std::string &Name, std::vector<std::string> Args) 148 : Name(Name), Args(std::move(Args)) {} 149 Function *codegen(); 150 const std::string &getName() const { return Name; } 151 }; 152 153 /// FunctionAST - This class represents a function definition itself. 154 class FunctionAST { 155 std::unique_ptr<PrototypeAST> Proto; 156 std::unique_ptr<ExprAST> Body; 157 158 public: 159 FunctionAST(std::unique_ptr<PrototypeAST> Proto, 160 std::unique_ptr<ExprAST> Body) 161 : Proto(std::move(Proto)), Body(std::move(Body)) {} 162 Function *codegen(); 163 }; 164 } // end anonymous namespace 165 166 //===----------------------------------------------------------------------===// 167 // Parser 168 //===----------------------------------------------------------------------===// 169 170 /// CurTok/getNextToken - Provide a simple token buffer. CurTok is the current 171 /// token the parser is looking at. getNextToken reads another token from the 172 /// lexer and updates CurTok with its results. 173 static int CurTok; 174 static int getNextToken() { return CurTok = gettok(); } 175 176 /// BinopPrecedence - This holds the precedence for each binary operator that is 177 /// defined. 178 static std::map<char, int> BinopPrecedence; 179 180 /// GetTokPrecedence - Get the precedence of the pending binary operator token. 181 static int GetTokPrecedence() { 182 if (!isascii(CurTok)) 183 return -1; 184 185 // Make sure it's a declared binop. 186 int TokPrec = BinopPrecedence[CurTok]; 187 if (TokPrec <= 0) 188 return -1; 189 return TokPrec; 190 } 191 192 /// LogError* - These are little helper functions for error handling. 193 std::unique_ptr<ExprAST> LogError(const char *Str) { 194 fprintf(stderr, "Error: %s\n", Str); 195 return nullptr; 196 } 197 198 std::unique_ptr<PrototypeAST> LogErrorP(const char *Str) { 199 LogError(Str); 200 return nullptr; 201 } 202 203 static std::unique_ptr<ExprAST> ParseExpression(); 204 205 /// numberexpr ::= number 206 static std::unique_ptr<ExprAST> ParseNumberExpr() { 207 auto Result = llvm::make_unique<NumberExprAST>(NumVal); 208 getNextToken(); // consume the number 209 return std::move(Result); 210 } 211 212 /// parenexpr ::= '(' expression ')' 213 static std::unique_ptr<ExprAST> ParseParenExpr() { 214 getNextToken(); // eat (. 215 auto V = ParseExpression(); 216 if (!V) 217 return nullptr; 218 219 if (CurTok != ')') 220 return LogError("expected ')'"); 221 getNextToken(); // eat ). 222 return V; 223 } 224 225 /// identifierexpr 226 /// ::= identifier 227 /// ::= identifier '(' expression* ')' 228 static std::unique_ptr<ExprAST> ParseIdentifierExpr() { 229 std::string IdName = IdentifierStr; 230 231 getNextToken(); // eat identifier. 232 233 if (CurTok != '(') // Simple variable ref. 234 return llvm::make_unique<VariableExprAST>(IdName); 235 236 // Call. 237 getNextToken(); // eat ( 238 std::vector<std::unique_ptr<ExprAST>> Args; 239 if (CurTok != ')') { 240 while (1) { 241 if (auto Arg = ParseExpression()) 242 Args.push_back(std::move(Arg)); 243 else 244 return nullptr; 245 246 if (CurTok == ')') 247 break; 248 249 if (CurTok != ',') 250 return LogError("Expected ')' or ',' in argument list"); 251 getNextToken(); 252 } 253 } 254 255 // Eat the ')'. 256 getNextToken(); 257 258 return llvm::make_unique<CallExprAST>(IdName, std::move(Args)); 259 } 260 261 /// primary 262 /// ::= identifierexpr 263 /// ::= numberexpr 264 /// ::= parenexpr 265 static std::unique_ptr<ExprAST> ParsePrimary() { 266 switch (CurTok) { 267 default: 268 return LogError("unknown token when expecting an expression"); 269 case tok_identifier: 270 return ParseIdentifierExpr(); 271 case tok_number: 272 return ParseNumberExpr(); 273 case '(': 274 return ParseParenExpr(); 275 } 276 } 277 278 /// binoprhs 279 /// ::= ('+' primary)* 280 static std::unique_ptr<ExprAST> ParseBinOpRHS(int ExprPrec, 281 std::unique_ptr<ExprAST> LHS) { 282 // If this is a binop, find its precedence. 283 while (1) { 284 int TokPrec = GetTokPrecedence(); 285 286 // If this is a binop that binds at least as tightly as the current binop, 287 // consume it, otherwise we are done. 288 if (TokPrec < ExprPrec) 289 return LHS; 290 291 // Okay, we know this is a binop. 292 int BinOp = CurTok; 293 getNextToken(); // eat binop 294 295 // Parse the primary expression after the binary operator. 296 auto RHS = ParsePrimary(); 297 if (!RHS) 298 return nullptr; 299 300 // If BinOp binds less tightly with RHS than the operator after RHS, let 301 // the pending operator take RHS as its LHS. 302 int NextPrec = GetTokPrecedence(); 303 if (TokPrec < NextPrec) { 304 RHS = ParseBinOpRHS(TokPrec + 1, std::move(RHS)); 305 if (!RHS) 306 return nullptr; 307 } 308 309 // Merge LHS/RHS. 310 LHS = 311 llvm::make_unique<BinaryExprAST>(BinOp, std::move(LHS), std::move(RHS)); 312 } 313 } 314 315 /// expression 316 /// ::= primary binoprhs 317 /// 318 static std::unique_ptr<ExprAST> ParseExpression() { 319 auto LHS = ParsePrimary(); 320 if (!LHS) 321 return nullptr; 322 323 return ParseBinOpRHS(0, std::move(LHS)); 324 } 325 326 /// prototype 327 /// ::= id '(' id* ')' 328 static std::unique_ptr<PrototypeAST> ParsePrototype() { 329 if (CurTok != tok_identifier) 330 return LogErrorP("Expected function name in prototype"); 331 332 std::string FnName = IdentifierStr; 333 getNextToken(); 334 335 if (CurTok != '(') 336 return LogErrorP("Expected '(' in prototype"); 337 338 std::vector<std::string> ArgNames; 339 while (getNextToken() == tok_identifier) 340 ArgNames.push_back(IdentifierStr); 341 if (CurTok != ')') 342 return LogErrorP("Expected ')' in prototype"); 343 344 // success. 345 getNextToken(); // eat ')'. 346 347 return llvm::make_unique<PrototypeAST>(FnName, std::move(ArgNames)); 348 } 349 350 /// definition ::= 'def' prototype expression 351 static std::unique_ptr<FunctionAST> ParseDefinition() { 352 getNextToken(); // eat def. 353 auto Proto = ParsePrototype(); 354 if (!Proto) 355 return nullptr; 356 357 if (auto E = ParseExpression()) 358 return llvm::make_unique<FunctionAST>(std::move(Proto), std::move(E)); 359 return nullptr; 360 } 361 362 /// toplevelexpr ::= expression 363 static std::unique_ptr<FunctionAST> ParseTopLevelExpr() { 364 if (auto E = ParseExpression()) { 365 // Make an anonymous proto. 366 auto Proto = llvm::make_unique<PrototypeAST>("__anon_expr", 367 std::vector<std::string>()); 368 return llvm::make_unique<FunctionAST>(std::move(Proto), std::move(E)); 369 } 370 return nullptr; 371 } 372 373 /// external ::= 'extern' prototype 374 static std::unique_ptr<PrototypeAST> ParseExtern() { 375 getNextToken(); // eat extern. 376 return ParsePrototype(); 377 } 378 379 //===----------------------------------------------------------------------===// 380 // Code Generation 381 //===----------------------------------------------------------------------===// 382 383 static std::unique_ptr<Module> TheModule; 384 static LLVMContext TheContext; 385 static IRBuilder<> Builder(TheContext); 386 static std::map<std::string, Value *> NamedValues; 387 388 Value *LogErrorV(const char *Str) { 389 LogError(Str); 390 return nullptr; 391 } 392 393 Value *NumberExprAST::codegen() { 394 return ConstantFP::get(TheContext, APFloat(Val)); 395 } 396 397 Value *VariableExprAST::codegen() { 398 // Look this variable up in the function. 399 Value *V = NamedValues[Name]; 400 if (!V) 401 return LogErrorV("Unknown variable name"); 402 return V; 403 } 404 405 Value *BinaryExprAST::codegen() { 406 Value *L = LHS->codegen(); 407 Value *R = RHS->codegen(); 408 if (!L || !R) 409 return nullptr; 410 411 switch (Op) { 412 case '+': 413 return Builder.CreateFAdd(L, R, "addtmp"); 414 case '-': 415 return Builder.CreateFSub(L, R, "subtmp"); 416 case '*': 417 return Builder.CreateFMul(L, R, "multmp"); 418 case '<': 419 L = Builder.CreateFCmpULT(L, R, "cmptmp"); 420 // Convert bool 0/1 to double 0.0 or 1.0 421 return Builder.CreateUIToFP(L, Type::getDoubleTy(TheContext), "booltmp"); 422 default: 423 return LogErrorV("invalid binary operator"); 424 } 425 } 426 427 Value *CallExprAST::codegen() { 428 // Look up the name in the global module table. 429 Function *CalleeF = TheModule->getFunction(Callee); 430 if (!CalleeF) 431 return LogErrorV("Unknown function referenced"); 432 433 // If argument mismatch error. 434 if (CalleeF->arg_size() != Args.size()) 435 return LogErrorV("Incorrect # arguments passed"); 436 437 std::vector<Value *> ArgsV; 438 for (unsigned i = 0, e = Args.size(); i != e; ++i) { 439 ArgsV.push_back(Args[i]->codegen()); 440 if (!ArgsV.back()) 441 return nullptr; 442 } 443 444 return Builder.CreateCall(CalleeF, ArgsV, "calltmp"); 445 } 446 447 Function *PrototypeAST::codegen() { 448 // Make the function type: double(double,double) etc. 449 std::vector<Type *> Doubles(Args.size(), Type::getDoubleTy(TheContext)); 450 FunctionType *FT = 451 FunctionType::get(Type::getDoubleTy(TheContext), Doubles, false); 452 453 Function *F = 454 Function::Create(FT, Function::ExternalLinkage, Name, TheModule.get()); 455 456 // Set names for all arguments. 457 unsigned Idx = 0; 458 for (auto &Arg : F->args()) 459 Arg.setName(Args[Idx++]); 460 461 return F; 462 } 463 464 Function *FunctionAST::codegen() { 465 // First, check for an existing function from a previous 'extern' declaration. 466 Function *TheFunction = TheModule->getFunction(Proto->getName()); 467 468 if (!TheFunction) 469 TheFunction = Proto->codegen(); 470 471 if (!TheFunction) 472 return nullptr; 473 474 // Create a new basic block to start insertion into. 475 BasicBlock *BB = BasicBlock::Create(TheContext, "entry", TheFunction); 476 Builder.SetInsertPoint(BB); 477 478 // Record the function arguments in the NamedValues map. 479 NamedValues.clear(); 480 for (auto &Arg : TheFunction->args()) 481 NamedValues[Arg.getName()] = &Arg; 482 483 if (Value *RetVal = Body->codegen()) { 484 // Finish off the function. 485 Builder.CreateRet(RetVal); 486 487 // Validate the generated code, checking for consistency. 488 verifyFunction(*TheFunction); 489 490 return TheFunction; 491 } 492 493 // Error reading body, remove function. 494 TheFunction->eraseFromParent(); 495 return nullptr; 496 } 497 498 //===----------------------------------------------------------------------===// 499 // Top-Level parsing and JIT Driver 500 //===----------------------------------------------------------------------===// 501 502 static void HandleDefinition() { 503 if (auto FnAST = ParseDefinition()) { 504 if (auto *FnIR = FnAST->codegen()) { 505 fprintf(stderr, "Read function definition:"); 506 FnIR->dump(); 507 } 508 } else { 509 // Skip token for error recovery. 510 getNextToken(); 511 } 512 } 513 514 static void HandleExtern() { 515 if (auto ProtoAST = ParseExtern()) { 516 if (auto *FnIR = ProtoAST->codegen()) { 517 fprintf(stderr, "Read extern: "); 518 FnIR->dump(); 519 } 520 } else { 521 // Skip token for error recovery. 522 getNextToken(); 523 } 524 } 525 526 static void HandleTopLevelExpression() { 527 // Evaluate a top-level expression into an anonymous function. 528 if (auto FnAST = ParseTopLevelExpr()) { 529 if (auto *FnIR = FnAST->codegen()) { 530 fprintf(stderr, "Read top-level expression:"); 531 FnIR->dump(); 532 } 533 } else { 534 // Skip token for error recovery. 535 getNextToken(); 536 } 537 } 538 539 /// top ::= definition | external | expression | ';' 540 static void MainLoop() { 541 while (1) { 542 fprintf(stderr, "ready> "); 543 switch (CurTok) { 544 case tok_eof: 545 return; 546 case ';': // ignore top-level semicolons. 547 getNextToken(); 548 break; 549 case tok_def: 550 HandleDefinition(); 551 break; 552 case tok_extern: 553 HandleExtern(); 554 break; 555 default: 556 HandleTopLevelExpression(); 557 break; 558 } 559 } 560 } 561 562 //===----------------------------------------------------------------------===// 563 // Main driver code. 564 //===----------------------------------------------------------------------===// 565 566 int main() { 567 // Install standard binary operators. 568 // 1 is lowest precedence. 569 BinopPrecedence['<'] = 10; 570 BinopPrecedence['+'] = 20; 571 BinopPrecedence['-'] = 20; 572 BinopPrecedence['*'] = 40; // highest. 573 574 // Prime the first token. 575 fprintf(stderr, "ready> "); 576 getNextToken(); 577 578 // Make the module, which holds all the code. 579 TheModule = llvm::make_unique<Module>("my cool jit", TheContext); 580 581 // Run the main "interpreter loop" now. 582 MainLoop(); 583 584 // Print out all of the generated code. 585 TheModule->dump(); 586 587 return 0; 588 } 589