1 #include "llvm/ADT/STLExtras.h" 2 #include "llvm/Analysis/Passes.h" 3 #include "llvm/IR/IRBuilder.h" 4 #include "llvm/IR/LLVMContext.h" 5 #include "llvm/IR/LegacyPassManager.h" 6 #include "llvm/IR/Module.h" 7 #include "llvm/IR/Verifier.h" 8 #include "llvm/Support/TargetSelect.h" 9 #include "llvm/Transforms/Scalar.h" 10 #include "llvm/Transforms/Scalar/GVN.h" 11 #include <cctype> 12 #include <cstdio> 13 #include <map> 14 #include <string> 15 #include <vector> 16 #include "../include/KaleidoscopeJIT.h" 17 18 using namespace llvm; 19 using namespace llvm::orc; 20 21 //===----------------------------------------------------------------------===// 22 // Lexer 23 //===----------------------------------------------------------------------===// 24 25 // The lexer returns tokens [0-255] if it is an unknown character, otherwise one 26 // of these for known things. 27 enum Token { 28 tok_eof = -1, 29 30 // commands 31 tok_def = -2, 32 tok_extern = -3, 33 34 // primary 35 tok_identifier = -4, 36 tok_number = -5, 37 38 // control 39 tok_if = -6, 40 tok_then = -7, 41 tok_else = -8, 42 tok_for = -9, 43 tok_in = -10 44 }; 45 46 static std::string IdentifierStr; // Filled in if tok_identifier 47 static double NumVal; // Filled in if tok_number 48 49 /// gettok - Return the next token from standard input. 50 static int gettok() { 51 static int LastChar = ' '; 52 53 // Skip any whitespace. 54 while (isspace(LastChar)) 55 LastChar = getchar(); 56 57 if (isalpha(LastChar)) { // identifier: [a-zA-Z][a-zA-Z0-9]* 58 IdentifierStr = LastChar; 59 while (isalnum((LastChar = getchar()))) 60 IdentifierStr += LastChar; 61 62 if (IdentifierStr == "def") 63 return tok_def; 64 if (IdentifierStr == "extern") 65 return tok_extern; 66 if (IdentifierStr == "if") 67 return tok_if; 68 if (IdentifierStr == "then") 69 return tok_then; 70 if (IdentifierStr == "else") 71 return tok_else; 72 if (IdentifierStr == "for") 73 return tok_for; 74 if (IdentifierStr == "in") 75 return tok_in; 76 return tok_identifier; 77 } 78 79 if (isdigit(LastChar) || LastChar == '.') { // Number: [0-9.]+ 80 std::string NumStr; 81 do { 82 NumStr += LastChar; 83 LastChar = getchar(); 84 } while (isdigit(LastChar) || LastChar == '.'); 85 86 NumVal = strtod(NumStr.c_str(), nullptr); 87 return tok_number; 88 } 89 90 if (LastChar == '#') { 91 // Comment until end of line. 92 do 93 LastChar = getchar(); 94 while (LastChar != EOF && LastChar != '\n' && LastChar != '\r'); 95 96 if (LastChar != EOF) 97 return gettok(); 98 } 99 100 // Check for end of file. Don't eat the EOF. 101 if (LastChar == EOF) 102 return tok_eof; 103 104 // Otherwise, just return the character as its ascii value. 105 int ThisChar = LastChar; 106 LastChar = getchar(); 107 return ThisChar; 108 } 109 110 //===----------------------------------------------------------------------===// 111 // Abstract Syntax Tree (aka Parse Tree) 112 //===----------------------------------------------------------------------===// 113 namespace { 114 /// ExprAST - Base class for all expression nodes. 115 class ExprAST { 116 public: 117 virtual ~ExprAST() {} 118 virtual Value *codegen() = 0; 119 }; 120 121 /// NumberExprAST - Expression class for numeric literals like "1.0". 122 class NumberExprAST : public ExprAST { 123 double Val; 124 125 public: 126 NumberExprAST(double Val) : Val(Val) {} 127 Value *codegen() override; 128 }; 129 130 /// VariableExprAST - Expression class for referencing a variable, like "a". 131 class VariableExprAST : public ExprAST { 132 std::string Name; 133 134 public: 135 VariableExprAST(const std::string &Name) : Name(Name) {} 136 Value *codegen() override; 137 }; 138 139 /// BinaryExprAST - Expression class for a binary operator. 140 class BinaryExprAST : public ExprAST { 141 char Op; 142 std::unique_ptr<ExprAST> LHS, RHS; 143 144 public: 145 BinaryExprAST(char Op, std::unique_ptr<ExprAST> LHS, 146 std::unique_ptr<ExprAST> RHS) 147 : Op(Op), LHS(std::move(LHS)), RHS(std::move(RHS)) {} 148 Value *codegen() override; 149 }; 150 151 /// CallExprAST - Expression class for function calls. 152 class CallExprAST : public ExprAST { 153 std::string Callee; 154 std::vector<std::unique_ptr<ExprAST>> Args; 155 156 public: 157 CallExprAST(const std::string &Callee, 158 std::vector<std::unique_ptr<ExprAST>> Args) 159 : Callee(Callee), Args(std::move(Args)) {} 160 Value *codegen() override; 161 }; 162 163 /// IfExprAST - Expression class for if/then/else. 164 class IfExprAST : public ExprAST { 165 std::unique_ptr<ExprAST> Cond, Then, Else; 166 167 public: 168 IfExprAST(std::unique_ptr<ExprAST> Cond, std::unique_ptr<ExprAST> Then, 169 std::unique_ptr<ExprAST> Else) 170 : Cond(std::move(Cond)), Then(std::move(Then)), Else(std::move(Else)) {} 171 Value *codegen() override; 172 }; 173 174 /// ForExprAST - Expression class for for/in. 175 class ForExprAST : public ExprAST { 176 std::string VarName; 177 std::unique_ptr<ExprAST> Start, End, Step, Body; 178 179 public: 180 ForExprAST(const std::string &VarName, std::unique_ptr<ExprAST> Start, 181 std::unique_ptr<ExprAST> End, std::unique_ptr<ExprAST> Step, 182 std::unique_ptr<ExprAST> Body) 183 : VarName(VarName), Start(std::move(Start)), End(std::move(End)), 184 Step(std::move(Step)), Body(std::move(Body)) {} 185 Value *codegen() override; 186 }; 187 188 /// PrototypeAST - This class represents the "prototype" for a function, 189 /// which captures its name, and its argument names (thus implicitly the number 190 /// of arguments the function takes). 191 class PrototypeAST { 192 std::string Name; 193 std::vector<std::string> Args; 194 195 public: 196 PrototypeAST(const std::string &Name, std::vector<std::string> Args) 197 : Name(Name), Args(std::move(Args)) {} 198 Function *codegen(); 199 const std::string &getName() const { return Name; } 200 }; 201 202 /// FunctionAST - This class represents a function definition itself. 203 class FunctionAST { 204 std::unique_ptr<PrototypeAST> Proto; 205 std::unique_ptr<ExprAST> Body; 206 207 public: 208 FunctionAST(std::unique_ptr<PrototypeAST> Proto, 209 std::unique_ptr<ExprAST> Body) 210 : Proto(std::move(Proto)), Body(std::move(Body)) {} 211 Function *codegen(); 212 }; 213 } // end anonymous namespace 214 215 //===----------------------------------------------------------------------===// 216 // Parser 217 //===----------------------------------------------------------------------===// 218 219 /// CurTok/getNextToken - Provide a simple token buffer. CurTok is the current 220 /// token the parser is looking at. getNextToken reads another token from the 221 /// lexer and updates CurTok with its results. 222 static int CurTok; 223 static int getNextToken() { return CurTok = gettok(); } 224 225 /// BinopPrecedence - This holds the precedence for each binary operator that is 226 /// defined. 227 static std::map<char, int> BinopPrecedence; 228 229 /// GetTokPrecedence - Get the precedence of the pending binary operator token. 230 static int GetTokPrecedence() { 231 if (!isascii(CurTok)) 232 return -1; 233 234 // Make sure it's a declared binop. 235 int TokPrec = BinopPrecedence[CurTok]; 236 if (TokPrec <= 0) 237 return -1; 238 return TokPrec; 239 } 240 241 /// LogError* - These are little helper functions for error handling. 242 std::unique_ptr<ExprAST> LogError(const char *Str) { 243 fprintf(stderr, "Error: %s\n", Str); 244 return nullptr; 245 } 246 247 std::unique_ptr<PrototypeAST> LogErrorP(const char *Str) { 248 LogError(Str); 249 return nullptr; 250 } 251 252 static std::unique_ptr<ExprAST> ParseExpression(); 253 254 /// numberexpr ::= number 255 static std::unique_ptr<ExprAST> ParseNumberExpr() { 256 auto Result = llvm::make_unique<NumberExprAST>(NumVal); 257 getNextToken(); // consume the number 258 return std::move(Result); 259 } 260 261 /// parenexpr ::= '(' expression ')' 262 static std::unique_ptr<ExprAST> ParseParenExpr() { 263 getNextToken(); // eat (. 264 auto V = ParseExpression(); 265 if (!V) 266 return nullptr; 267 268 if (CurTok != ')') 269 return LogError("expected ')'"); 270 getNextToken(); // eat ). 271 return V; 272 } 273 274 /// identifierexpr 275 /// ::= identifier 276 /// ::= identifier '(' expression* ')' 277 static std::unique_ptr<ExprAST> ParseIdentifierExpr() { 278 std::string IdName = IdentifierStr; 279 280 getNextToken(); // eat identifier. 281 282 if (CurTok != '(') // Simple variable ref. 283 return llvm::make_unique<VariableExprAST>(IdName); 284 285 // Call. 286 getNextToken(); // eat ( 287 std::vector<std::unique_ptr<ExprAST>> Args; 288 if (CurTok != ')') { 289 while (1) { 290 if (auto Arg = ParseExpression()) 291 Args.push_back(std::move(Arg)); 292 else 293 return nullptr; 294 295 if (CurTok == ')') 296 break; 297 298 if (CurTok != ',') 299 return LogError("Expected ')' or ',' in argument list"); 300 getNextToken(); 301 } 302 } 303 304 // Eat the ')'. 305 getNextToken(); 306 307 return llvm::make_unique<CallExprAST>(IdName, std::move(Args)); 308 } 309 310 /// ifexpr ::= 'if' expression 'then' expression 'else' expression 311 static std::unique_ptr<ExprAST> ParseIfExpr() { 312 getNextToken(); // eat the if. 313 314 // condition. 315 auto Cond = ParseExpression(); 316 if (!Cond) 317 return nullptr; 318 319 if (CurTok != tok_then) 320 return LogError("expected then"); 321 getNextToken(); // eat the then 322 323 auto Then = ParseExpression(); 324 if (!Then) 325 return nullptr; 326 327 if (CurTok != tok_else) 328 return LogError("expected else"); 329 330 getNextToken(); 331 332 auto Else = ParseExpression(); 333 if (!Else) 334 return nullptr; 335 336 return llvm::make_unique<IfExprAST>(std::move(Cond), std::move(Then), 337 std::move(Else)); 338 } 339 340 /// forexpr ::= 'for' identifier '=' expr ',' expr (',' expr)? 'in' expression 341 static std::unique_ptr<ExprAST> ParseForExpr() { 342 getNextToken(); // eat the for. 343 344 if (CurTok != tok_identifier) 345 return LogError("expected identifier after for"); 346 347 std::string IdName = IdentifierStr; 348 getNextToken(); // eat identifier. 349 350 if (CurTok != '=') 351 return LogError("expected '=' after for"); 352 getNextToken(); // eat '='. 353 354 auto Start = ParseExpression(); 355 if (!Start) 356 return nullptr; 357 if (CurTok != ',') 358 return LogError("expected ',' after for start value"); 359 getNextToken(); 360 361 auto End = ParseExpression(); 362 if (!End) 363 return nullptr; 364 365 // The step value is optional. 366 std::unique_ptr<ExprAST> Step; 367 if (CurTok == ',') { 368 getNextToken(); 369 Step = ParseExpression(); 370 if (!Step) 371 return nullptr; 372 } 373 374 if (CurTok != tok_in) 375 return LogError("expected 'in' after for"); 376 getNextToken(); // eat 'in'. 377 378 auto Body = ParseExpression(); 379 if (!Body) 380 return nullptr; 381 382 return llvm::make_unique<ForExprAST>(IdName, std::move(Start), std::move(End), 383 std::move(Step), std::move(Body)); 384 } 385 386 /// primary 387 /// ::= identifierexpr 388 /// ::= numberexpr 389 /// ::= parenexpr 390 /// ::= ifexpr 391 /// ::= forexpr 392 static std::unique_ptr<ExprAST> ParsePrimary() { 393 switch (CurTok) { 394 default: 395 return LogError("unknown token when expecting an expression"); 396 case tok_identifier: 397 return ParseIdentifierExpr(); 398 case tok_number: 399 return ParseNumberExpr(); 400 case '(': 401 return ParseParenExpr(); 402 case tok_if: 403 return ParseIfExpr(); 404 case tok_for: 405 return ParseForExpr(); 406 } 407 } 408 409 /// binoprhs 410 /// ::= ('+' primary)* 411 static std::unique_ptr<ExprAST> ParseBinOpRHS(int ExprPrec, 412 std::unique_ptr<ExprAST> LHS) { 413 // If this is a binop, find its precedence. 414 while (1) { 415 int TokPrec = GetTokPrecedence(); 416 417 // If this is a binop that binds at least as tightly as the current binop, 418 // consume it, otherwise we are done. 419 if (TokPrec < ExprPrec) 420 return LHS; 421 422 // Okay, we know this is a binop. 423 int BinOp = CurTok; 424 getNextToken(); // eat binop 425 426 // Parse the primary expression after the binary operator. 427 auto RHS = ParsePrimary(); 428 if (!RHS) 429 return nullptr; 430 431 // If BinOp binds less tightly with RHS than the operator after RHS, let 432 // the pending operator take RHS as its LHS. 433 int NextPrec = GetTokPrecedence(); 434 if (TokPrec < NextPrec) { 435 RHS = ParseBinOpRHS(TokPrec + 1, std::move(RHS)); 436 if (!RHS) 437 return nullptr; 438 } 439 440 // Merge LHS/RHS. 441 LHS = 442 llvm::make_unique<BinaryExprAST>(BinOp, std::move(LHS), std::move(RHS)); 443 } 444 } 445 446 /// expression 447 /// ::= primary binoprhs 448 /// 449 static std::unique_ptr<ExprAST> ParseExpression() { 450 auto LHS = ParsePrimary(); 451 if (!LHS) 452 return nullptr; 453 454 return ParseBinOpRHS(0, std::move(LHS)); 455 } 456 457 /// prototype 458 /// ::= id '(' id* ')' 459 static std::unique_ptr<PrototypeAST> ParsePrototype() { 460 if (CurTok != tok_identifier) 461 return LogErrorP("Expected function name in prototype"); 462 463 std::string FnName = IdentifierStr; 464 getNextToken(); 465 466 if (CurTok != '(') 467 return LogErrorP("Expected '(' in prototype"); 468 469 std::vector<std::string> ArgNames; 470 while (getNextToken() == tok_identifier) 471 ArgNames.push_back(IdentifierStr); 472 if (CurTok != ')') 473 return LogErrorP("Expected ')' in prototype"); 474 475 // success. 476 getNextToken(); // eat ')'. 477 478 return llvm::make_unique<PrototypeAST>(FnName, std::move(ArgNames)); 479 } 480 481 /// definition ::= 'def' prototype expression 482 static std::unique_ptr<FunctionAST> ParseDefinition() { 483 getNextToken(); // eat def. 484 auto Proto = ParsePrototype(); 485 if (!Proto) 486 return nullptr; 487 488 if (auto E = ParseExpression()) 489 return llvm::make_unique<FunctionAST>(std::move(Proto), std::move(E)); 490 return nullptr; 491 } 492 493 /// toplevelexpr ::= expression 494 static std::unique_ptr<FunctionAST> ParseTopLevelExpr() { 495 if (auto E = ParseExpression()) { 496 // Make an anonymous proto. 497 auto Proto = llvm::make_unique<PrototypeAST>("__anon_expr", 498 std::vector<std::string>()); 499 return llvm::make_unique<FunctionAST>(std::move(Proto), std::move(E)); 500 } 501 return nullptr; 502 } 503 504 /// external ::= 'extern' prototype 505 static std::unique_ptr<PrototypeAST> ParseExtern() { 506 getNextToken(); // eat extern. 507 return ParsePrototype(); 508 } 509 510 //===----------------------------------------------------------------------===// 511 // Code Generation 512 //===----------------------------------------------------------------------===// 513 514 static std::unique_ptr<Module> TheModule; 515 static IRBuilder<> Builder(getGlobalContext()); 516 static std::map<std::string, Value *> NamedValues; 517 static std::unique_ptr<legacy::FunctionPassManager> TheFPM; 518 static std::unique_ptr<KaleidoscopeJIT> TheJIT; 519 static std::map<std::string, std::unique_ptr<PrototypeAST>> FunctionProtos; 520 521 Value *LogErrorV(const char *Str) { 522 LogError(Str); 523 return nullptr; 524 } 525 526 Function *getFunction(std::string Name) { 527 // First, see if the function has already been added to the current module. 528 if (auto *F = TheModule->getFunction(Name)) 529 return F; 530 531 // If not, check whether we can codegen the declaration from some existing 532 // prototype. 533 auto FI = FunctionProtos.find(Name); 534 if (FI != FunctionProtos.end()) 535 return FI->second->codegen(); 536 537 // If no existing prototype exists, return null. 538 return nullptr; 539 } 540 541 Value *NumberExprAST::codegen() { 542 return ConstantFP::get(getGlobalContext(), APFloat(Val)); 543 } 544 545 Value *VariableExprAST::codegen() { 546 // Look this variable up in the function. 547 Value *V = NamedValues[Name]; 548 if (!V) 549 return LogErrorV("Unknown variable name"); 550 return V; 551 } 552 553 Value *BinaryExprAST::codegen() { 554 Value *L = LHS->codegen(); 555 Value *R = RHS->codegen(); 556 if (!L || !R) 557 return nullptr; 558 559 switch (Op) { 560 case '+': 561 return Builder.CreateFAdd(L, R, "addtmp"); 562 case '-': 563 return Builder.CreateFSub(L, R, "subtmp"); 564 case '*': 565 return Builder.CreateFMul(L, R, "multmp"); 566 case '<': 567 L = Builder.CreateFCmpULT(L, R, "cmptmp"); 568 // Convert bool 0/1 to double 0.0 or 1.0 569 return Builder.CreateUIToFP(L, Type::getDoubleTy(getGlobalContext()), 570 "booltmp"); 571 default: 572 return LogErrorV("invalid binary operator"); 573 } 574 } 575 576 Value *CallExprAST::codegen() { 577 // Look up the name in the global module table. 578 Function *CalleeF = getFunction(Callee); 579 if (!CalleeF) 580 return LogErrorV("Unknown function referenced"); 581 582 // If argument mismatch error. 583 if (CalleeF->arg_size() != Args.size()) 584 return LogErrorV("Incorrect # arguments passed"); 585 586 std::vector<Value *> ArgsV; 587 for (unsigned i = 0, e = Args.size(); i != e; ++i) { 588 ArgsV.push_back(Args[i]->codegen()); 589 if (!ArgsV.back()) 590 return nullptr; 591 } 592 593 return Builder.CreateCall(CalleeF, ArgsV, "calltmp"); 594 } 595 596 Value *IfExprAST::codegen() { 597 Value *CondV = Cond->codegen(); 598 if (!CondV) 599 return nullptr; 600 601 // Convert condition to a bool by comparing equal to 0.0. 602 CondV = Builder.CreateFCmpONE( 603 CondV, ConstantFP::get(getGlobalContext(), APFloat(0.0)), "ifcond"); 604 605 Function *TheFunction = Builder.GetInsertBlock()->getParent(); 606 607 // Create blocks for the then and else cases. Insert the 'then' block at the 608 // end of the function. 609 BasicBlock *ThenBB = 610 BasicBlock::Create(getGlobalContext(), "then", TheFunction); 611 BasicBlock *ElseBB = BasicBlock::Create(getGlobalContext(), "else"); 612 BasicBlock *MergeBB = BasicBlock::Create(getGlobalContext(), "ifcont"); 613 614 Builder.CreateCondBr(CondV, ThenBB, ElseBB); 615 616 // Emit then value. 617 Builder.SetInsertPoint(ThenBB); 618 619 Value *ThenV = Then->codegen(); 620 if (!ThenV) 621 return nullptr; 622 623 Builder.CreateBr(MergeBB); 624 // Codegen of 'Then' can change the current block, update ThenBB for the PHI. 625 ThenBB = Builder.GetInsertBlock(); 626 627 // Emit else block. 628 TheFunction->getBasicBlockList().push_back(ElseBB); 629 Builder.SetInsertPoint(ElseBB); 630 631 Value *ElseV = Else->codegen(); 632 if (!ElseV) 633 return nullptr; 634 635 Builder.CreateBr(MergeBB); 636 // Codegen of 'Else' can change the current block, update ElseBB for the PHI. 637 ElseBB = Builder.GetInsertBlock(); 638 639 // Emit merge block. 640 TheFunction->getBasicBlockList().push_back(MergeBB); 641 Builder.SetInsertPoint(MergeBB); 642 PHINode *PN = 643 Builder.CreatePHI(Type::getDoubleTy(getGlobalContext()), 2, "iftmp"); 644 645 PN->addIncoming(ThenV, ThenBB); 646 PN->addIncoming(ElseV, ElseBB); 647 return PN; 648 } 649 650 // Output for-loop as: 651 // ... 652 // start = startexpr 653 // goto loop 654 // loop: 655 // variable = phi [start, loopheader], [nextvariable, loopend] 656 // ... 657 // bodyexpr 658 // ... 659 // loopend: 660 // step = stepexpr 661 // nextvariable = variable + step 662 // endcond = endexpr 663 // br endcond, loop, endloop 664 // outloop: 665 Value *ForExprAST::codegen() { 666 // Emit the start code first, without 'variable' in scope. 667 Value *StartVal = Start->codegen(); 668 if (!StartVal) 669 return nullptr; 670 671 // Make the new basic block for the loop header, inserting after current 672 // block. 673 Function *TheFunction = Builder.GetInsertBlock()->getParent(); 674 BasicBlock *PreheaderBB = Builder.GetInsertBlock(); 675 BasicBlock *LoopBB = 676 BasicBlock::Create(getGlobalContext(), "loop", TheFunction); 677 678 // Insert an explicit fall through from the current block to the LoopBB. 679 Builder.CreateBr(LoopBB); 680 681 // Start insertion in LoopBB. 682 Builder.SetInsertPoint(LoopBB); 683 684 // Start the PHI node with an entry for Start. 685 PHINode *Variable = Builder.CreatePHI(Type::getDoubleTy(getGlobalContext()), 686 2, VarName.c_str()); 687 Variable->addIncoming(StartVal, PreheaderBB); 688 689 // Within the loop, the variable is defined equal to the PHI node. If it 690 // shadows an existing variable, we have to restore it, so save it now. 691 Value *OldVal = NamedValues[VarName]; 692 NamedValues[VarName] = Variable; 693 694 // Emit the body of the loop. This, like any other expr, can change the 695 // current BB. Note that we ignore the value computed by the body, but don't 696 // allow an error. 697 if (!Body->codegen()) 698 return nullptr; 699 700 // Emit the step value. 701 Value *StepVal = nullptr; 702 if (Step) { 703 StepVal = Step->codegen(); 704 if (!StepVal) 705 return nullptr; 706 } else { 707 // If not specified, use 1.0. 708 StepVal = ConstantFP::get(getGlobalContext(), APFloat(1.0)); 709 } 710 711 Value *NextVar = Builder.CreateFAdd(Variable, StepVal, "nextvar"); 712 713 // Compute the end condition. 714 Value *EndCond = End->codegen(); 715 if (!EndCond) 716 return nullptr; 717 718 // Convert condition to a bool by comparing equal to 0.0. 719 EndCond = Builder.CreateFCmpONE( 720 EndCond, ConstantFP::get(getGlobalContext(), APFloat(0.0)), "loopcond"); 721 722 // Create the "after loop" block and insert it. 723 BasicBlock *LoopEndBB = Builder.GetInsertBlock(); 724 BasicBlock *AfterBB = 725 BasicBlock::Create(getGlobalContext(), "afterloop", TheFunction); 726 727 // Insert the conditional branch into the end of LoopEndBB. 728 Builder.CreateCondBr(EndCond, LoopBB, AfterBB); 729 730 // Any new code will be inserted in AfterBB. 731 Builder.SetInsertPoint(AfterBB); 732 733 // Add a new entry to the PHI node for the backedge. 734 Variable->addIncoming(NextVar, LoopEndBB); 735 736 // Restore the unshadowed variable. 737 if (OldVal) 738 NamedValues[VarName] = OldVal; 739 else 740 NamedValues.erase(VarName); 741 742 // for expr always returns 0.0. 743 return Constant::getNullValue(Type::getDoubleTy(getGlobalContext())); 744 } 745 746 Function *PrototypeAST::codegen() { 747 // Make the function type: double(double,double) etc. 748 std::vector<Type *> Doubles(Args.size(), 749 Type::getDoubleTy(getGlobalContext())); 750 FunctionType *FT = 751 FunctionType::get(Type::getDoubleTy(getGlobalContext()), Doubles, false); 752 753 Function *F = 754 Function::Create(FT, Function::ExternalLinkage, Name, TheModule.get()); 755 756 // Set names for all arguments. 757 unsigned Idx = 0; 758 for (auto &Arg : F->args()) 759 Arg.setName(Args[Idx++]); 760 761 return F; 762 } 763 764 Function *FunctionAST::codegen() { 765 // Transfer ownership of the prototype to the FunctionProtos map, but keep a 766 // reference to it for use below. 767 auto &P = *Proto; 768 FunctionProtos[Proto->getName()] = std::move(Proto); 769 Function *TheFunction = getFunction(P.getName()); 770 if (!TheFunction) 771 return nullptr; 772 773 // Create a new basic block to start insertion into. 774 BasicBlock *BB = BasicBlock::Create(getGlobalContext(), "entry", TheFunction); 775 Builder.SetInsertPoint(BB); 776 777 // Record the function arguments in the NamedValues map. 778 NamedValues.clear(); 779 for (auto &Arg : TheFunction->args()) 780 NamedValues[Arg.getName()] = &Arg; 781 782 if (Value *RetVal = Body->codegen()) { 783 // Finish off the function. 784 Builder.CreateRet(RetVal); 785 786 // Validate the generated code, checking for consistency. 787 verifyFunction(*TheFunction); 788 789 // Run the optimizer on the function. 790 TheFPM->run(*TheFunction); 791 792 return TheFunction; 793 } 794 795 // Error reading body, remove function. 796 TheFunction->eraseFromParent(); 797 return nullptr; 798 } 799 800 //===----------------------------------------------------------------------===// 801 // Top-Level parsing and JIT Driver 802 //===----------------------------------------------------------------------===// 803 804 static void InitializeModuleAndPassManager() { 805 // Open a new module. 806 TheModule = llvm::make_unique<Module>("my cool jit", getGlobalContext()); 807 TheModule->setDataLayout(TheJIT->getTargetMachine().createDataLayout()); 808 809 // Create a new pass manager attached to it. 810 TheFPM = llvm::make_unique<legacy::FunctionPassManager>(TheModule.get()); 811 812 // Do simple "peephole" optimizations and bit-twiddling optzns. 813 TheFPM->add(createInstructionCombiningPass()); 814 // Reassociate expressions. 815 TheFPM->add(createReassociatePass()); 816 // Eliminate Common SubExpressions. 817 TheFPM->add(createGVNPass()); 818 // Simplify the control flow graph (deleting unreachable blocks, etc). 819 TheFPM->add(createCFGSimplificationPass()); 820 821 TheFPM->doInitialization(); 822 } 823 824 static void HandleDefinition() { 825 if (auto FnAST = ParseDefinition()) { 826 if (auto *FnIR = FnAST->codegen()) { 827 fprintf(stderr, "Read function definition:"); 828 FnIR->dump(); 829 TheJIT->addModule(std::move(TheModule)); 830 InitializeModuleAndPassManager(); 831 } 832 } else { 833 // Skip token for error recovery. 834 getNextToken(); 835 } 836 } 837 838 static void HandleExtern() { 839 if (auto ProtoAST = ParseExtern()) { 840 if (auto *FnIR = ProtoAST->codegen()) { 841 fprintf(stderr, "Read extern: "); 842 FnIR->dump(); 843 FunctionProtos[ProtoAST->getName()] = std::move(ProtoAST); 844 } 845 } else { 846 // Skip token for error recovery. 847 getNextToken(); 848 } 849 } 850 851 static void HandleTopLevelExpression() { 852 // Evaluate a top-level expression into an anonymous function. 853 if (auto FnAST = ParseTopLevelExpr()) { 854 if (FnAST->codegen()) { 855 856 // JIT the module containing the anonymous expression, keeping a handle so 857 // we can free it later. 858 auto H = TheJIT->addModule(std::move(TheModule)); 859 InitializeModuleAndPassManager(); 860 861 // Search the JIT for the __anon_expr symbol. 862 auto ExprSymbol = TheJIT->findSymbol("__anon_expr"); 863 assert(ExprSymbol && "Function not found"); 864 865 // Get the symbol's address and cast it to the right type (takes no 866 // arguments, returns a double) so we can call it as a native function. 867 double (*FP)() = (double (*)())(intptr_t)ExprSymbol.getAddress(); 868 fprintf(stderr, "Evaluated to %f\n", FP()); 869 870 // Delete the anonymous expression module from the JIT. 871 TheJIT->removeModule(H); 872 } 873 } else { 874 // Skip token for error recovery. 875 getNextToken(); 876 } 877 } 878 879 /// top ::= definition | external | expression | ';' 880 static void MainLoop() { 881 while (1) { 882 fprintf(stderr, "ready> "); 883 switch (CurTok) { 884 case tok_eof: 885 return; 886 case ';': // ignore top-level semicolons. 887 getNextToken(); 888 break; 889 case tok_def: 890 HandleDefinition(); 891 break; 892 case tok_extern: 893 HandleExtern(); 894 break; 895 default: 896 HandleTopLevelExpression(); 897 break; 898 } 899 } 900 } 901 902 //===----------------------------------------------------------------------===// 903 // "Library" functions that can be "extern'd" from user code. 904 //===----------------------------------------------------------------------===// 905 906 /// putchard - putchar that takes a double and returns 0. 907 extern "C" double putchard(double X) { 908 fputc((char)X, stderr); 909 return 0; 910 } 911 912 /// printd - printf that takes a double prints it as "%f\n", returning 0. 913 extern "C" double printd(double X) { 914 fprintf(stderr, "%f\n", X); 915 return 0; 916 } 917 918 //===----------------------------------------------------------------------===// 919 // Main driver code. 920 //===----------------------------------------------------------------------===// 921 922 int main() { 923 InitializeNativeTarget(); 924 InitializeNativeTargetAsmPrinter(); 925 InitializeNativeTargetAsmParser(); 926 927 // Install standard binary operators. 928 // 1 is lowest precedence. 929 BinopPrecedence['<'] = 10; 930 BinopPrecedence['+'] = 20; 931 BinopPrecedence['-'] = 20; 932 BinopPrecedence['*'] = 40; // highest. 933 934 // Prime the first token. 935 fprintf(stderr, "ready> "); 936 getNextToken(); 937 938 TheJIT = llvm::make_unique<KaleidoscopeJIT>(); 939 940 InitializeModuleAndPassManager(); 941 942 // Run the main "interpreter loop" now. 943 MainLoop(); 944 945 return 0; 946 } 947