1 //===- OpFormatGen.cpp - MLIR operation asm format generator --------------===// 2 // 3 // Part of the LLVM Project, under the Apache License v2.0 with LLVM Exceptions. 4 // See https://llvm.org/LICENSE.txt for license information. 5 // SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception 6 // 7 //===----------------------------------------------------------------------===// 8 9 #include "OpFormatGen.h" 10 #include "mlir/Support/LogicalResult.h" 11 #include "mlir/TableGen/Format.h" 12 #include "mlir/TableGen/GenInfo.h" 13 #include "mlir/TableGen/OpClass.h" 14 #include "mlir/TableGen/OpInterfaces.h" 15 #include "mlir/TableGen/OpTrait.h" 16 #include "mlir/TableGen/Operator.h" 17 #include "llvm/ADT/MapVector.h" 18 #include "llvm/ADT/Sequence.h" 19 #include "llvm/ADT/SmallBitVector.h" 20 #include "llvm/ADT/StringExtras.h" 21 #include "llvm/ADT/TypeSwitch.h" 22 #include "llvm/Support/CommandLine.h" 23 #include "llvm/Support/Signals.h" 24 #include "llvm/TableGen/Error.h" 25 #include "llvm/TableGen/Record.h" 26 27 #define DEBUG_TYPE "mlir-tblgen-opformatgen" 28 29 using namespace mlir; 30 using namespace mlir::tblgen; 31 32 static llvm::cl::opt<bool> formatErrorIsFatal( 33 "asmformat-error-is-fatal", 34 llvm::cl::desc("Emit a fatal error if format parsing fails"), 35 llvm::cl::init(true)); 36 37 //===----------------------------------------------------------------------===// 38 // Element 39 //===----------------------------------------------------------------------===// 40 41 namespace { 42 /// This class represents a single format element. 43 class Element { 44 public: 45 enum class Kind { 46 /// This element is a directive. 47 AttrDictDirective, 48 FunctionalTypeDirective, 49 OperandsDirective, 50 ResultsDirective, 51 SuccessorsDirective, 52 TypeDirective, 53 54 /// This element is a literal. 55 Literal, 56 57 /// This element is an variable value. 58 AttributeVariable, 59 OperandVariable, 60 ResultVariable, 61 SuccessorVariable, 62 63 /// This element is an optional element. 64 Optional, 65 }; 66 Element(Kind kind) : kind(kind) {} 67 virtual ~Element() = default; 68 69 /// Return the kind of this element. 70 Kind getKind() const { return kind; } 71 72 private: 73 /// The kind of this element. 74 Kind kind; 75 }; 76 } // namespace 77 78 //===----------------------------------------------------------------------===// 79 // VariableElement 80 81 namespace { 82 /// This class represents an instance of an variable element. A variable refers 83 /// to something registered on the operation itself, e.g. an argument, result, 84 /// etc. 85 template <typename VarT, Element::Kind kindVal> 86 class VariableElement : public Element { 87 public: 88 VariableElement(const VarT *var) : Element(kindVal), var(var) {} 89 static bool classof(const Element *element) { 90 return element->getKind() == kindVal; 91 } 92 const VarT *getVar() { return var; } 93 94 protected: 95 const VarT *var; 96 }; 97 98 /// This class represents a variable that refers to an attribute argument. 99 struct AttributeVariable 100 : public VariableElement<NamedAttribute, Element::Kind::AttributeVariable> { 101 using VariableElement<NamedAttribute, 102 Element::Kind::AttributeVariable>::VariableElement; 103 104 /// Return the constant builder call for the type of this attribute, or None 105 /// if it doesn't have one. 106 Optional<StringRef> getTypeBuilder() const { 107 Optional<Type> attrType = var->attr.getValueType(); 108 return attrType ? attrType->getBuilderCall() : llvm::None; 109 } 110 }; 111 112 /// This class represents a variable that refers to an operand argument. 113 using OperandVariable = 114 VariableElement<NamedTypeConstraint, Element::Kind::OperandVariable>; 115 116 /// This class represents a variable that refers to a result. 117 using ResultVariable = 118 VariableElement<NamedTypeConstraint, Element::Kind::ResultVariable>; 119 120 /// This class represents a variable that refers to a successor. 121 using SuccessorVariable = 122 VariableElement<NamedSuccessor, Element::Kind::SuccessorVariable>; 123 } // end anonymous namespace 124 125 //===----------------------------------------------------------------------===// 126 // DirectiveElement 127 128 namespace { 129 /// This class implements single kind directives. 130 template <Element::Kind type> 131 class DirectiveElement : public Element { 132 public: 133 DirectiveElement() : Element(type){}; 134 static bool classof(const Element *ele) { return ele->getKind() == type; } 135 }; 136 /// This class represents the `operands` directive. This directive represents 137 /// all of the operands of an operation. 138 using OperandsDirective = DirectiveElement<Element::Kind::OperandsDirective>; 139 140 /// This class represents the `results` directive. This directive represents 141 /// all of the results of an operation. 142 using ResultsDirective = DirectiveElement<Element::Kind::ResultsDirective>; 143 144 /// This class represents the `successors` directive. This directive represents 145 /// all of the successors of an operation. 146 using SuccessorsDirective = 147 DirectiveElement<Element::Kind::SuccessorsDirective>; 148 149 /// This class represents the `attr-dict` directive. This directive represents 150 /// the attribute dictionary of the operation. 151 class AttrDictDirective 152 : public DirectiveElement<Element::Kind::AttrDictDirective> { 153 public: 154 explicit AttrDictDirective(bool withKeyword) : withKeyword(withKeyword) {} 155 bool isWithKeyword() const { return withKeyword; } 156 157 private: 158 /// If the dictionary should be printed with the 'attributes' keyword. 159 bool withKeyword; 160 }; 161 162 /// This class represents the `functional-type` directive. This directive takes 163 /// two arguments and formats them, respectively, as the inputs and results of a 164 /// FunctionType. 165 class FunctionalTypeDirective 166 : public DirectiveElement<Element::Kind::FunctionalTypeDirective> { 167 public: 168 FunctionalTypeDirective(std::unique_ptr<Element> inputs, 169 std::unique_ptr<Element> results) 170 : inputs(std::move(inputs)), results(std::move(results)) {} 171 Element *getInputs() const { return inputs.get(); } 172 Element *getResults() const { return results.get(); } 173 174 private: 175 /// The input and result arguments. 176 std::unique_ptr<Element> inputs, results; 177 }; 178 179 /// This class represents the `type` directive. 180 class TypeDirective : public DirectiveElement<Element::Kind::TypeDirective> { 181 public: 182 TypeDirective(std::unique_ptr<Element> arg) : operand(std::move(arg)) {} 183 Element *getOperand() const { return operand.get(); } 184 185 private: 186 /// The operand that is used to format the directive. 187 std::unique_ptr<Element> operand; 188 }; 189 } // end anonymous namespace 190 191 //===----------------------------------------------------------------------===// 192 // LiteralElement 193 194 namespace { 195 /// This class represents an instance of a literal element. 196 class LiteralElement : public Element { 197 public: 198 LiteralElement(StringRef literal) 199 : Element{Kind::Literal}, literal(literal) {} 200 static bool classof(const Element *element) { 201 return element->getKind() == Kind::Literal; 202 } 203 204 /// Return the literal for this element. 205 StringRef getLiteral() const { return literal; } 206 207 /// Returns true if the given string is a valid literal. 208 static bool isValidLiteral(StringRef value); 209 210 private: 211 /// The spelling of the literal for this element. 212 StringRef literal; 213 }; 214 } // end anonymous namespace 215 216 bool LiteralElement::isValidLiteral(StringRef value) { 217 if (value.empty()) 218 return false; 219 char front = value.front(); 220 221 // If there is only one character, this must either be punctuation or a 222 // single character bare identifier. 223 if (value.size() == 1) 224 return isalpha(front) || StringRef("_:,=<>()[]").contains(front); 225 226 // Check the punctuation that are larger than a single character. 227 if (value == "->") 228 return true; 229 230 // Otherwise, this must be an identifier. 231 if (!isalpha(front) && front != '_') 232 return false; 233 return llvm::all_of(value.drop_front(), [](char c) { 234 return isalnum(c) || c == '_' || c == '$' || c == '.'; 235 }); 236 } 237 238 //===----------------------------------------------------------------------===// 239 // OptionalElement 240 241 namespace { 242 /// This class represents a group of elements that are optionally emitted based 243 /// upon an optional variable of the operation. 244 class OptionalElement : public Element { 245 public: 246 OptionalElement(std::vector<std::unique_ptr<Element>> &&elements, 247 unsigned anchor) 248 : Element{Kind::Optional}, elements(std::move(elements)), anchor(anchor) { 249 } 250 static bool classof(const Element *element) { 251 return element->getKind() == Kind::Optional; 252 } 253 254 /// Return the nested elements of this grouping. 255 auto getElements() const { return llvm::make_pointee_range(elements); } 256 257 /// Return the anchor of this optional group. 258 Element *getAnchor() const { return elements[anchor].get(); } 259 260 private: 261 /// The child elements of this optional. 262 std::vector<std::unique_ptr<Element>> elements; 263 /// The index of the element that acts as the anchor for the optional group. 264 unsigned anchor; 265 }; 266 } // end anonymous namespace 267 268 //===----------------------------------------------------------------------===// 269 // OperationFormat 270 //===----------------------------------------------------------------------===// 271 272 namespace { 273 struct OperationFormat { 274 /// This class represents a specific resolver for an operand or result type. 275 class TypeResolution { 276 public: 277 TypeResolution() = default; 278 279 /// Get the index into the buildable types for this type, or None. 280 Optional<int> getBuilderIdx() const { return builderIdx; } 281 void setBuilderIdx(int idx) { builderIdx = idx; } 282 283 /// Get the variable this type is resolved to, or None. 284 const NamedTypeConstraint *getVariable() const { return variable; } 285 Optional<StringRef> getVarTransformer() const { 286 return variableTransformer; 287 } 288 void setVariable(const NamedTypeConstraint *var, 289 Optional<StringRef> transformer) { 290 variable = var; 291 variableTransformer = transformer; 292 } 293 294 private: 295 /// If the type is resolved with a buildable type, this is the index into 296 /// 'buildableTypes' in the parent format. 297 Optional<int> builderIdx; 298 /// If the type is resolved based upon another operand or result, this is 299 /// the variable that this type is resolved to. 300 const NamedTypeConstraint *variable; 301 /// If the type is resolved based upon another operand or result, this is 302 /// a transformer to apply to the variable when resolving. 303 Optional<StringRef> variableTransformer; 304 }; 305 306 OperationFormat(const Operator &op) 307 : allOperands(false), allOperandTypes(false), allResultTypes(false) { 308 operandTypes.resize(op.getNumOperands(), TypeResolution()); 309 resultTypes.resize(op.getNumResults(), TypeResolution()); 310 } 311 312 /// Generate the operation parser from this format. 313 void genParser(Operator &op, OpClass &opClass); 314 /// Generate the c++ to resolve the types of operands and results during 315 /// parsing. 316 void genParserTypeResolution(Operator &op, OpMethodBody &body); 317 /// Generate the c++ to resolve successors during parsing. 318 void genParserSuccessorResolution(Operator &op, OpMethodBody &body); 319 /// Generate the c++ to handling variadic segment size traits. 320 void genParserVariadicSegmentResolution(Operator &op, OpMethodBody &body); 321 322 /// Generate the operation printer from this format. 323 void genPrinter(Operator &op, OpClass &opClass); 324 325 /// The various elements in this format. 326 std::vector<std::unique_ptr<Element>> elements; 327 328 /// A flag indicating if all operand/result types were seen. If the format 329 /// contains these, it can not contain individual type resolvers. 330 bool allOperands, allOperandTypes, allResultTypes; 331 332 /// A map of buildable types to indices. 333 llvm::MapVector<StringRef, int, llvm::StringMap<int>> buildableTypes; 334 335 /// The index of the buildable type, if valid, for every operand and result. 336 std::vector<TypeResolution> operandTypes, resultTypes; 337 }; 338 } // end anonymous namespace 339 340 //===----------------------------------------------------------------------===// 341 // Parser Gen 342 343 /// Returns if we can format the given attribute as an EnumAttr in the parser 344 /// format. 345 static bool canFormatEnumAttr(const NamedAttribute *attr) { 346 const EnumAttr *enumAttr = dyn_cast<EnumAttr>(&attr->attr); 347 if (!enumAttr) 348 return false; 349 350 // The attribute must have a valid underlying type and a constant builder. 351 return !enumAttr->getUnderlyingType().empty() && 352 !enumAttr->getConstBuilderTemplate().empty(); 353 } 354 355 /// The code snippet used to generate a parser call for an attribute. 356 /// 357 /// {0}: The storage type of the attribute. 358 /// {1}: The name of the attribute. 359 /// {2}: The type for the attribute. 360 const char *const attrParserCode = R"( 361 {0} {1}Attr; 362 if (parser.parseAttribute({1}Attr{2}, "{1}", result.attributes)) 363 return failure(); 364 )"; 365 366 /// The code snippet used to generate a parser call for an enum attribute. 367 /// 368 /// {0}: The name of the attribute. 369 /// {1}: The c++ namespace for the enum symbolize functions. 370 /// {2}: The function to symbolize a string of the enum. 371 /// {3}: The constant builder call to create an attribute of the enum type. 372 const char *const enumAttrParserCode = R"( 373 { 374 StringAttr attrVal; 375 SmallVector<NamedAttribute, 1> attrStorage; 376 auto loc = parser.getCurrentLocation(); 377 if (parser.parseAttribute(attrVal, parser.getBuilder().getNoneType(), 378 "{0}", attrStorage)) 379 return failure(); 380 381 auto attrOptional = {1}::{2}(attrVal.getValue()); 382 if (!attrOptional) 383 return parser.emitError(loc, "invalid ") 384 << "{0} attribute specification: " << attrVal; 385 386 result.addAttribute("{0}", {3}); 387 } 388 )"; 389 390 /// The code snippet used to generate a parser call for an operand. 391 /// 392 /// {0}: The name of the operand. 393 const char *const variadicOperandParserCode = R"( 394 if (parser.parseOperandList({0}Operands)) 395 return failure(); 396 )"; 397 const char *const optionalOperandParserCode = R"( 398 { 399 OpAsmParser::OperandType operand; 400 OptionalParseResult parseResult = parser.parseOptionalOperand(operand); 401 if (parseResult.hasValue()) { 402 if (failed(*parseResult)) 403 return failure(); 404 {0}Operands.push_back(operand); 405 } 406 } 407 )"; 408 const char *const operandParserCode = R"( 409 if (parser.parseOperand({0}RawOperands[0])) 410 return failure(); 411 )"; 412 413 /// The code snippet used to generate a parser call for a type list. 414 /// 415 /// {0}: The name for the type list. 416 const char *const variadicTypeParserCode = R"( 417 if (parser.parseTypeList({0}Types)) 418 return failure(); 419 )"; 420 const char *const optionalTypeParserCode = R"( 421 { 422 Type optionalType; 423 OptionalParseResult parseResult = parser.parseOptionalType(optionalType); 424 if (parseResult.hasValue()) { 425 if (failed(*parseResult)) 426 return failure(); 427 {0}Types.push_back(optionalType); 428 } 429 } 430 )"; 431 const char *const typeParserCode = R"( 432 if (parser.parseType({0}RawTypes[0])) 433 return failure(); 434 )"; 435 436 /// The code snippet used to generate a parser call for a functional type. 437 /// 438 /// {0}: The name for the input type list. 439 /// {1}: The name for the result type list. 440 const char *const functionalTypeParserCode = R"( 441 FunctionType {0}__{1}_functionType; 442 if (parser.parseType({0}__{1}_functionType)) 443 return failure(); 444 {0}Types = {0}__{1}_functionType.getInputs(); 445 {1}Types = {0}__{1}_functionType.getResults(); 446 )"; 447 448 /// The code snippet used to generate a parser call for a successor list. 449 /// 450 /// {0}: The name for the successor list. 451 const char *successorListParserCode = R"( 452 SmallVector<Block *, 2> {0}Successors; 453 { 454 Block *succ; 455 auto firstSucc = parser.parseOptionalSuccessor(succ); 456 if (firstSucc.hasValue()) { 457 if (failed(*firstSucc)) 458 return failure(); 459 {0}Successors.emplace_back(succ); 460 461 // Parse any trailing successors. 462 while (succeeded(parser.parseOptionalComma())) { 463 if (parser.parseSuccessor(succ)) 464 return failure(); 465 {0}Successors.emplace_back(succ); 466 } 467 } 468 } 469 )"; 470 471 /// The code snippet used to generate a parser call for a successor. 472 /// 473 /// {0}: The name of the successor. 474 const char *successorParserCode = R"( 475 Block *{0}Successor = nullptr; 476 if (parser.parseSuccessor({0}Successor)) 477 return failure(); 478 )"; 479 480 namespace { 481 /// The type of length for a given parse argument. 482 enum class ArgumentLengthKind { 483 /// The argument is variadic, and may contain 0->N elements. 484 Variadic, 485 /// The argument is optional, and may contain 0 or 1 elements. 486 Optional, 487 /// The argument is a single element, i.e. always represents 1 element. 488 Single 489 }; 490 } // end anonymous namespace 491 492 /// Get the length kind for the given constraint. 493 static ArgumentLengthKind 494 getArgumentLengthKind(const NamedTypeConstraint *var) { 495 if (var->isOptional()) 496 return ArgumentLengthKind::Optional; 497 if (var->isVariadic()) 498 return ArgumentLengthKind::Variadic; 499 return ArgumentLengthKind::Single; 500 } 501 502 /// Get the name used for the type list for the given type directive operand. 503 /// 'lengthKind' to the corresponding kind for the given argument. 504 static StringRef getTypeListName(Element *arg, ArgumentLengthKind &lengthKind) { 505 if (auto *operand = dyn_cast<OperandVariable>(arg)) { 506 lengthKind = getArgumentLengthKind(operand->getVar()); 507 return operand->getVar()->name; 508 } 509 if (auto *result = dyn_cast<ResultVariable>(arg)) { 510 lengthKind = getArgumentLengthKind(result->getVar()); 511 return result->getVar()->name; 512 } 513 lengthKind = ArgumentLengthKind::Variadic; 514 if (isa<OperandsDirective>(arg)) 515 return "allOperand"; 516 if (isa<ResultsDirective>(arg)) 517 return "allResult"; 518 llvm_unreachable("unknown 'type' directive argument"); 519 } 520 521 /// Generate the parser for a literal value. 522 static void genLiteralParser(StringRef value, OpMethodBody &body) { 523 // Handle the case of a keyword/identifier. 524 if (value.front() == '_' || isalpha(value.front())) { 525 body << "Keyword(\"" << value << "\")"; 526 return; 527 } 528 body << (StringRef)llvm::StringSwitch<StringRef>(value) 529 .Case("->", "Arrow()") 530 .Case(":", "Colon()") 531 .Case(",", "Comma()") 532 .Case("=", "Equal()") 533 .Case("<", "Less()") 534 .Case(">", "Greater()") 535 .Case("(", "LParen()") 536 .Case(")", "RParen()") 537 .Case("[", "LSquare()") 538 .Case("]", "RSquare()"); 539 } 540 541 /// Generate the storage code required for parsing the given element. 542 static void genElementParserStorage(Element *element, OpMethodBody &body) { 543 if (auto *optional = dyn_cast<OptionalElement>(element)) { 544 for (auto &childElement : optional->getElements()) 545 genElementParserStorage(&childElement, body); 546 } else if (auto *operand = dyn_cast<OperandVariable>(element)) { 547 StringRef name = operand->getVar()->name; 548 if (operand->getVar()->isVariableLength()) { 549 body << " SmallVector<OpAsmParser::OperandType, 4> " << name 550 << "Operands;\n"; 551 } else { 552 body << " OpAsmParser::OperandType " << name << "RawOperands[1];\n" 553 << " ArrayRef<OpAsmParser::OperandType> " << name << "Operands(" 554 << name << "RawOperands);"; 555 } 556 body << llvm::formatv( 557 " llvm::SMLoc {0}OperandsLoc = parser.getCurrentLocation();\n" 558 " (void){0}OperandsLoc;\n", 559 name); 560 } else if (auto *dir = dyn_cast<TypeDirective>(element)) { 561 ArgumentLengthKind lengthKind; 562 StringRef name = getTypeListName(dir->getOperand(), lengthKind); 563 if (lengthKind != ArgumentLengthKind::Single) 564 body << " SmallVector<Type, 1> " << name << "Types;\n"; 565 else 566 body << llvm::formatv(" Type {0}RawTypes[1];\n", name) 567 << llvm::formatv(" ArrayRef<Type> {0}Types({0}RawTypes);\n", name); 568 } else if (auto *dir = dyn_cast<FunctionalTypeDirective>(element)) { 569 ArgumentLengthKind ignored; 570 body << " ArrayRef<Type> " << getTypeListName(dir->getInputs(), ignored) 571 << "Types;\n"; 572 body << " ArrayRef<Type> " << getTypeListName(dir->getResults(), ignored) 573 << "Types;\n"; 574 } 575 } 576 577 /// Generate the parser for a single format element. 578 static void genElementParser(Element *element, OpMethodBody &body, 579 FmtContext &attrTypeCtx) { 580 /// Optional Group. 581 if (auto *optional = dyn_cast<OptionalElement>(element)) { 582 auto elements = optional->getElements(); 583 584 // Generate a special optional parser for the first element to gate the 585 // parsing of the rest of the elements. 586 if (auto *literal = dyn_cast<LiteralElement>(&*elements.begin())) { 587 body << " if (succeeded(parser.parseOptional"; 588 genLiteralParser(literal->getLiteral(), body); 589 body << ")) {\n"; 590 } else if (auto *opVar = dyn_cast<OperandVariable>(&*elements.begin())) { 591 genElementParser(opVar, body, attrTypeCtx); 592 body << " if (!" << opVar->getVar()->name << "Operands.empty()) {\n"; 593 } 594 595 // Generate the rest of the elements normally. 596 for (auto &childElement : llvm::drop_begin(elements, 1)) 597 genElementParser(&childElement, body, attrTypeCtx); 598 body << " }\n"; 599 600 /// Literals. 601 } else if (LiteralElement *literal = dyn_cast<LiteralElement>(element)) { 602 body << " if (parser.parse"; 603 genLiteralParser(literal->getLiteral(), body); 604 body << ")\n return failure();\n"; 605 606 /// Arguments. 607 } else if (auto *attr = dyn_cast<AttributeVariable>(element)) { 608 const NamedAttribute *var = attr->getVar(); 609 610 // Check to see if we can parse this as an enum attribute. 611 if (canFormatEnumAttr(var)) { 612 const EnumAttr &enumAttr = cast<EnumAttr>(var->attr); 613 614 // Generate the code for building an attribute for this enum. 615 std::string attrBuilderStr; 616 { 617 llvm::raw_string_ostream os(attrBuilderStr); 618 os << tgfmt(enumAttr.getConstBuilderTemplate(), &attrTypeCtx, 619 "attrOptional.getValue()"); 620 } 621 622 body << formatv(enumAttrParserCode, var->name, enumAttr.getCppNamespace(), 623 enumAttr.getStringToSymbolFnName(), attrBuilderStr); 624 return; 625 } 626 627 // If this attribute has a buildable type, use that when parsing the 628 // attribute. 629 std::string attrTypeStr; 630 if (Optional<StringRef> typeBuilder = attr->getTypeBuilder()) { 631 llvm::raw_string_ostream os(attrTypeStr); 632 os << ", " << tgfmt(*typeBuilder, &attrTypeCtx); 633 } 634 635 body << formatv(attrParserCode, var->attr.getStorageType(), var->name, 636 attrTypeStr); 637 } else if (auto *operand = dyn_cast<OperandVariable>(element)) { 638 ArgumentLengthKind lengthKind = getArgumentLengthKind(operand->getVar()); 639 StringRef name = operand->getVar()->name; 640 if (lengthKind == ArgumentLengthKind::Variadic) 641 body << llvm::formatv(variadicOperandParserCode, name); 642 else if (lengthKind == ArgumentLengthKind::Optional) 643 body << llvm::formatv(optionalOperandParserCode, name); 644 else 645 body << formatv(operandParserCode, name); 646 } else if (auto *successor = dyn_cast<SuccessorVariable>(element)) { 647 bool isVariadic = successor->getVar()->isVariadic(); 648 body << formatv(isVariadic ? successorListParserCode : successorParserCode, 649 successor->getVar()->name); 650 651 /// Directives. 652 } else if (auto *attrDict = dyn_cast<AttrDictDirective>(element)) { 653 body << " if (parser.parseOptionalAttrDict" 654 << (attrDict->isWithKeyword() ? "WithKeyword" : "") 655 << "(result.attributes))\n" 656 << " return failure();\n"; 657 } else if (isa<OperandsDirective>(element)) { 658 body << " llvm::SMLoc allOperandLoc = parser.getCurrentLocation();\n" 659 << " SmallVector<OpAsmParser::OperandType, 4> allOperands;\n" 660 << " if (parser.parseOperandList(allOperands))\n" 661 << " return failure();\n"; 662 } else if (isa<SuccessorsDirective>(element)) { 663 body << llvm::formatv(successorListParserCode, "full"); 664 } else if (auto *dir = dyn_cast<TypeDirective>(element)) { 665 ArgumentLengthKind lengthKind; 666 StringRef listName = getTypeListName(dir->getOperand(), lengthKind); 667 if (lengthKind == ArgumentLengthKind::Variadic) 668 body << llvm::formatv(variadicTypeParserCode, listName); 669 else if (lengthKind == ArgumentLengthKind::Optional) 670 body << llvm::formatv(optionalTypeParserCode, listName); 671 else 672 body << formatv(typeParserCode, listName); 673 } else if (auto *dir = dyn_cast<FunctionalTypeDirective>(element)) { 674 ArgumentLengthKind ignored; 675 body << formatv(functionalTypeParserCode, 676 getTypeListName(dir->getInputs(), ignored), 677 getTypeListName(dir->getResults(), ignored)); 678 } else { 679 llvm_unreachable("unknown format element"); 680 } 681 } 682 683 void OperationFormat::genParser(Operator &op, OpClass &opClass) { 684 auto &method = opClass.newMethod( 685 "ParseResult", "parse", "OpAsmParser &parser, OperationState &result", 686 OpMethod::MP_Static); 687 auto &body = method.body(); 688 689 // Generate variables to store the operands and type within the format. This 690 // allows for referencing these variables in the presence of optional 691 // groupings. 692 for (auto &element : elements) 693 genElementParserStorage(&*element, body); 694 695 // A format context used when parsing attributes with buildable types. 696 FmtContext attrTypeCtx; 697 attrTypeCtx.withBuilder("parser.getBuilder()"); 698 699 // Generate parsers for each of the elements. 700 for (auto &element : elements) 701 genElementParser(element.get(), body, attrTypeCtx); 702 703 // Generate the code to resolve the operand/result types and successors now 704 // that they have been parsed. 705 genParserTypeResolution(op, body); 706 genParserSuccessorResolution(op, body); 707 genParserVariadicSegmentResolution(op, body); 708 709 // Mark the operation as having resizable operand list if required. 710 if (op.hasResizableOperandList()) 711 body << " result.setOperandListToResizable();\n"; 712 713 body << " return success();\n"; 714 } 715 716 void OperationFormat::genParserTypeResolution(Operator &op, 717 OpMethodBody &body) { 718 // If any of type resolutions use transformed variables, make sure that the 719 // types of those variables are resolved. 720 SmallPtrSet<const NamedTypeConstraint *, 8> verifiedVariables; 721 FmtContext verifierFCtx; 722 for (TypeResolution &resolver : 723 llvm::concat<TypeResolution>(resultTypes, operandTypes)) { 724 Optional<StringRef> transformer = resolver.getVarTransformer(); 725 if (!transformer) 726 continue; 727 // Ensure that we don't verify the same variables twice. 728 const NamedTypeConstraint *variable = resolver.getVariable(); 729 if (!verifiedVariables.insert(variable).second) 730 continue; 731 732 auto constraint = variable->constraint; 733 body << " for (Type type : " << variable->name << "Types) {\n" 734 << " (void)type;\n" 735 << " if (!(" 736 << tgfmt(constraint.getConditionTemplate(), 737 &verifierFCtx.withSelf("type")) 738 << ")) {\n" 739 << formatv(" return parser.emitError(parser.getNameLoc()) << " 740 "\"'{0}' must be {1}, but got \" << type;\n", 741 variable->name, constraint.getDescription()) 742 << " }\n" 743 << " }\n"; 744 } 745 746 // Initialize the set of buildable types. 747 if (!buildableTypes.empty()) { 748 body << " Builder &builder = parser.getBuilder();\n"; 749 750 FmtContext typeBuilderCtx; 751 typeBuilderCtx.withBuilder("builder"); 752 for (auto &it : buildableTypes) 753 body << " Type odsBuildableType" << it.second << " = " 754 << tgfmt(it.first, &typeBuilderCtx) << ";\n"; 755 } 756 757 // Emit the code necessary for a type resolver. 758 auto emitTypeResolver = [&](TypeResolution &resolver, StringRef curVar) { 759 if (Optional<int> val = resolver.getBuilderIdx()) { 760 body << "odsBuildableType" << *val; 761 } else if (const NamedTypeConstraint *var = resolver.getVariable()) { 762 if (Optional<StringRef> tform = resolver.getVarTransformer()) 763 body << tgfmt(*tform, &FmtContext().withSelf(var->name + "Types[0]")); 764 else 765 body << var->name << "Types"; 766 } else { 767 body << curVar << "Types"; 768 } 769 }; 770 771 // Resolve each of the result types. 772 if (allResultTypes) { 773 body << " result.addTypes(allResultTypes);\n"; 774 } else { 775 for (unsigned i = 0, e = op.getNumResults(); i != e; ++i) { 776 body << " result.addTypes("; 777 emitTypeResolver(resultTypes[i], op.getResultName(i)); 778 body << ");\n"; 779 } 780 } 781 782 // Early exit if there are no operands. 783 if (op.getNumOperands() == 0) 784 return; 785 786 // Handle the case where all operand types are in one group. 787 if (allOperandTypes) { 788 // If we have all operands together, use the full operand list directly. 789 if (allOperands) { 790 body << " if (parser.resolveOperands(allOperands, allOperandTypes, " 791 "allOperandLoc, result.operands))\n" 792 " return failure();\n"; 793 return; 794 } 795 796 // Otherwise, use llvm::concat to merge the disjoint operand lists together. 797 // llvm::concat does not allow the case of a single range, so guard it here. 798 body << " if (parser.resolveOperands("; 799 if (op.getNumOperands() > 1) { 800 body << "llvm::concat<const OpAsmParser::OperandType>("; 801 llvm::interleaveComma(op.getOperands(), body, [&](auto &operand) { 802 body << operand.name << "Operands"; 803 }); 804 body << ")"; 805 } else { 806 body << op.operand_begin()->name << "Operands"; 807 } 808 body << ", allOperandTypes, parser.getNameLoc(), result.operands))\n" 809 << " return failure();\n"; 810 return; 811 } 812 // Handle the case where all of the operands were grouped together. 813 if (allOperands) { 814 body << " if (parser.resolveOperands(allOperands, "; 815 816 // Group all of the operand types together to perform the resolution all at 817 // once. Use llvm::concat to perform the merge. llvm::concat does not allow 818 // the case of a single range, so guard it here. 819 if (op.getNumOperands() > 1) { 820 body << "llvm::concat<const Type>("; 821 llvm::interleaveComma( 822 llvm::seq<int>(0, op.getNumOperands()), body, [&](int i) { 823 body << "ArrayRef<Type>("; 824 emitTypeResolver(operandTypes[i], op.getOperand(i).name); 825 body << ")"; 826 }); 827 body << ")"; 828 } else { 829 emitTypeResolver(operandTypes.front(), op.getOperand(0).name); 830 } 831 832 body << ", allOperandLoc, result.operands))\n" 833 << " return failure();\n"; 834 return; 835 } 836 837 // The final case is the one where each of the operands types are resolved 838 // separately. 839 for (unsigned i = 0, e = op.getNumOperands(); i != e; ++i) { 840 NamedTypeConstraint &operand = op.getOperand(i); 841 body << " if (parser.resolveOperands(" << operand.name << "Operands, "; 842 emitTypeResolver(operandTypes[i], operand.name); 843 844 // If this isn't a buildable type, verify the sizes match by adding the loc. 845 if (!operandTypes[i].getBuilderIdx()) 846 body << ", " << operand.name << "OperandsLoc"; 847 body << ", result.operands))\n return failure();\n"; 848 } 849 } 850 851 void OperationFormat::genParserSuccessorResolution(Operator &op, 852 OpMethodBody &body) { 853 // Check for the case where all successors were parsed. 854 bool hasAllSuccessors = llvm::any_of( 855 elements, [](auto &elt) { return isa<SuccessorsDirective>(elt.get()); }); 856 if (hasAllSuccessors) { 857 body << " result.addSuccessors(fullSuccessors);\n"; 858 return; 859 } 860 861 // Otherwise, handle each successor individually. 862 for (const NamedSuccessor &successor : op.getSuccessors()) { 863 if (successor.isVariadic()) 864 body << " result.addSuccessors(" << successor.name << "Successors);\n"; 865 else 866 body << " result.addSuccessors(" << successor.name << "Successor);\n"; 867 } 868 } 869 870 void OperationFormat::genParserVariadicSegmentResolution(Operator &op, 871 OpMethodBody &body) { 872 if (!allOperands && op.getTrait("OpTrait::AttrSizedOperandSegments")) { 873 body << " result.addAttribute(\"operand_segment_sizes\", " 874 << "builder.getI32VectorAttr({"; 875 auto interleaveFn = [&](const NamedTypeConstraint &operand) { 876 // If the operand is variadic emit the parsed size. 877 if (operand.isVariableLength()) 878 body << "static_cast<int32_t>(" << operand.name << "Operands.size())"; 879 else 880 body << "1"; 881 }; 882 llvm::interleaveComma(op.getOperands(), body, interleaveFn); 883 body << "}));\n"; 884 } 885 } 886 887 //===----------------------------------------------------------------------===// 888 // PrinterGen 889 890 /// Generate the printer for the 'attr-dict' directive. 891 static void genAttrDictPrinter(OperationFormat &fmt, Operator &op, 892 OpMethodBody &body, bool withKeyword) { 893 // Collect all of the attributes used in the format, these will be elided. 894 SmallVector<const NamedAttribute *, 1> usedAttributes; 895 for (auto &it : fmt.elements) 896 if (auto *attr = dyn_cast<AttributeVariable>(it.get())) 897 usedAttributes.push_back(attr->getVar()); 898 899 body << " p.printOptionalAttrDict" << (withKeyword ? "WithKeyword" : "") 900 << "(getAttrs(), /*elidedAttrs=*/{"; 901 // Elide the variadic segment size attributes if necessary. 902 if (!fmt.allOperands && op.getTrait("OpTrait::AttrSizedOperandSegments")) 903 body << "\"operand_segment_sizes\", "; 904 llvm::interleaveComma(usedAttributes, body, [&](const NamedAttribute *attr) { 905 body << "\"" << attr->name << "\""; 906 }); 907 body << "});\n"; 908 } 909 910 /// Generate the printer for a literal value. `shouldEmitSpace` is true if a 911 /// space should be emitted before this element. `lastWasPunctuation` is true if 912 /// the previous element was a punctuation literal. 913 static void genLiteralPrinter(StringRef value, OpMethodBody &body, 914 bool &shouldEmitSpace, bool &lastWasPunctuation) { 915 body << " p"; 916 917 // Don't insert a space for certain punctuation. 918 auto shouldPrintSpaceBeforeLiteral = [&] { 919 if (value.size() != 1 && value != "->") 920 return true; 921 if (lastWasPunctuation) 922 return !StringRef(">)}],").contains(value.front()); 923 return !StringRef("<>(){}[],").contains(value.front()); 924 }; 925 if (shouldEmitSpace && shouldPrintSpaceBeforeLiteral()) 926 body << " << \" \""; 927 body << " << \"" << value << "\";\n"; 928 929 // Insert a space after certain literals. 930 shouldEmitSpace = 931 value.size() != 1 || !StringRef("<({[").contains(value.front()); 932 lastWasPunctuation = !(value.front() == '_' || isalpha(value.front())); 933 } 934 935 /// Generate the C++ for an operand to a (*-)type directive. 936 static OpMethodBody &genTypeOperandPrinter(Element *arg, OpMethodBody &body) { 937 if (isa<OperandsDirective>(arg)) 938 return body << "getOperation()->getOperandTypes()"; 939 if (isa<ResultsDirective>(arg)) 940 return body << "getOperation()->getResultTypes()"; 941 auto *operand = dyn_cast<OperandVariable>(arg); 942 auto *var = operand ? operand->getVar() : cast<ResultVariable>(arg)->getVar(); 943 if (var->isVariadic()) 944 return body << var->name << "().getTypes()"; 945 if (var->isOptional()) 946 return body << llvm::formatv( 947 "({0}() ? ArrayRef<Type>({0}().getType()) : ArrayRef<Type>())", 948 var->name); 949 return body << "ArrayRef<Type>(" << var->name << "().getType())"; 950 } 951 952 /// Generate the code for printing the given element. 953 static void genElementPrinter(Element *element, OpMethodBody &body, 954 OperationFormat &fmt, Operator &op, 955 bool &shouldEmitSpace, bool &lastWasPunctuation) { 956 if (LiteralElement *literal = dyn_cast<LiteralElement>(element)) 957 return genLiteralPrinter(literal->getLiteral(), body, shouldEmitSpace, 958 lastWasPunctuation); 959 960 // Emit an optional group. 961 if (OptionalElement *optional = dyn_cast<OptionalElement>(element)) { 962 // Emit the check for the presence of the anchor element. 963 Element *anchor = optional->getAnchor(); 964 if (auto *operand = dyn_cast<OperandVariable>(anchor)) { 965 const NamedTypeConstraint *var = operand->getVar(); 966 if (var->isOptional()) 967 body << " if (" << var->name << "()) {\n"; 968 else if (var->isVariadic()) 969 body << " if (!" << var->name << "().empty()) {\n"; 970 } else { 971 body << " if (getAttr(\"" 972 << cast<AttributeVariable>(anchor)->getVar()->name << "\")) {\n"; 973 } 974 975 // Emit each of the elements. 976 for (Element &childElement : optional->getElements()) 977 genElementPrinter(&childElement, body, fmt, op, shouldEmitSpace, 978 lastWasPunctuation); 979 body << " }\n"; 980 return; 981 } 982 983 // Emit the attribute dictionary. 984 if (auto *attrDict = dyn_cast<AttrDictDirective>(element)) { 985 genAttrDictPrinter(fmt, op, body, attrDict->isWithKeyword()); 986 lastWasPunctuation = false; 987 return; 988 } 989 990 // Optionally insert a space before the next element. The AttrDict printer 991 // already adds a space as necessary. 992 if (shouldEmitSpace || !lastWasPunctuation) 993 body << " p << \" \";\n"; 994 lastWasPunctuation = false; 995 shouldEmitSpace = true; 996 997 if (auto *attr = dyn_cast<AttributeVariable>(element)) { 998 const NamedAttribute *var = attr->getVar(); 999 1000 // If we are formatting as an enum, symbolize the attribute as a string. 1001 if (canFormatEnumAttr(var)) { 1002 const EnumAttr &enumAttr = cast<EnumAttr>(var->attr); 1003 body << " p << \"\\\"\" << " << enumAttr.getSymbolToStringFnName() << "(" 1004 << var->name << "()) << \"\\\"\";\n"; 1005 return; 1006 } 1007 1008 // Elide the attribute type if it is buildable. 1009 if (attr->getTypeBuilder()) 1010 body << " p.printAttributeWithoutType(" << var->name << "Attr());\n"; 1011 else 1012 body << " p.printAttribute(" << var->name << "Attr());\n"; 1013 } else if (auto *operand = dyn_cast<OperandVariable>(element)) { 1014 if (operand->getVar()->isOptional()) { 1015 body << " if (Value value = " << operand->getVar()->name << "())\n" 1016 << " p << value;\n"; 1017 } else { 1018 body << " p << " << operand->getVar()->name << "();\n"; 1019 } 1020 } else if (auto *successor = dyn_cast<SuccessorVariable>(element)) { 1021 const NamedSuccessor *var = successor->getVar(); 1022 if (var->isVariadic()) 1023 body << " llvm::interleaveComma(" << var->name << "(), p);\n"; 1024 else 1025 body << " p << " << var->name << "();\n"; 1026 } else if (isa<OperandsDirective>(element)) { 1027 body << " p << getOperation()->getOperands();\n"; 1028 } else if (isa<SuccessorsDirective>(element)) { 1029 body << " llvm::interleaveComma(getOperation()->getSuccessors(), p);\n"; 1030 } else if (auto *dir = dyn_cast<TypeDirective>(element)) { 1031 body << " p << "; 1032 genTypeOperandPrinter(dir->getOperand(), body) << ";\n"; 1033 } else if (auto *dir = dyn_cast<FunctionalTypeDirective>(element)) { 1034 body << " p.printFunctionalType("; 1035 genTypeOperandPrinter(dir->getInputs(), body) << ", "; 1036 genTypeOperandPrinter(dir->getResults(), body) << ");\n"; 1037 } else { 1038 llvm_unreachable("unknown format element"); 1039 } 1040 } 1041 1042 void OperationFormat::genPrinter(Operator &op, OpClass &opClass) { 1043 auto &method = opClass.newMethod("void", "print", "OpAsmPrinter &p"); 1044 auto &body = method.body(); 1045 1046 // Emit the operation name, trimming the prefix if this is the standard 1047 // dialect. 1048 body << " p << \""; 1049 std::string opName = op.getOperationName(); 1050 if (op.getDialectName() == "std") 1051 body << StringRef(opName).drop_front(4); 1052 else 1053 body << opName; 1054 body << "\";\n"; 1055 1056 // Flags for if we should emit a space, and if the last element was 1057 // punctuation. 1058 bool shouldEmitSpace = true, lastWasPunctuation = false; 1059 for (auto &element : elements) 1060 genElementPrinter(element.get(), body, *this, op, shouldEmitSpace, 1061 lastWasPunctuation); 1062 } 1063 1064 //===----------------------------------------------------------------------===// 1065 // FormatLexer 1066 //===----------------------------------------------------------------------===// 1067 1068 namespace { 1069 /// This class represents a specific token in the input format. 1070 class Token { 1071 public: 1072 enum Kind { 1073 // Markers. 1074 eof, 1075 error, 1076 1077 // Tokens with no info. 1078 l_paren, 1079 r_paren, 1080 caret, 1081 comma, 1082 equal, 1083 question, 1084 1085 // Keywords. 1086 keyword_start, 1087 kw_attr_dict, 1088 kw_attr_dict_w_keyword, 1089 kw_functional_type, 1090 kw_operands, 1091 kw_results, 1092 kw_successors, 1093 kw_type, 1094 keyword_end, 1095 1096 // String valued tokens. 1097 identifier, 1098 literal, 1099 variable, 1100 }; 1101 Token(Kind kind, StringRef spelling) : kind(kind), spelling(spelling) {} 1102 1103 /// Return the bytes that make up this token. 1104 StringRef getSpelling() const { return spelling; } 1105 1106 /// Return the kind of this token. 1107 Kind getKind() const { return kind; } 1108 1109 /// Return a location for this token. 1110 llvm::SMLoc getLoc() const { 1111 return llvm::SMLoc::getFromPointer(spelling.data()); 1112 } 1113 1114 /// Return if this token is a keyword. 1115 bool isKeyword() const { return kind > keyword_start && kind < keyword_end; } 1116 1117 private: 1118 /// Discriminator that indicates the kind of token this is. 1119 Kind kind; 1120 1121 /// A reference to the entire token contents; this is always a pointer into 1122 /// a memory buffer owned by the source manager. 1123 StringRef spelling; 1124 }; 1125 1126 /// This class implements a simple lexer for operation assembly format strings. 1127 class FormatLexer { 1128 public: 1129 FormatLexer(llvm::SourceMgr &mgr, Operator &op); 1130 1131 /// Lex the next token and return it. 1132 Token lexToken(); 1133 1134 /// Emit an error to the lexer with the given location and message. 1135 Token emitError(llvm::SMLoc loc, const Twine &msg); 1136 Token emitError(const char *loc, const Twine &msg); 1137 1138 Token emitErrorAndNote(llvm::SMLoc loc, const Twine &msg, const Twine ¬e); 1139 1140 private: 1141 Token formToken(Token::Kind kind, const char *tokStart) { 1142 return Token(kind, StringRef(tokStart, curPtr - tokStart)); 1143 } 1144 1145 /// Return the next character in the stream. 1146 int getNextChar(); 1147 1148 /// Lex an identifier, literal, or variable. 1149 Token lexIdentifier(const char *tokStart); 1150 Token lexLiteral(const char *tokStart); 1151 Token lexVariable(const char *tokStart); 1152 1153 llvm::SourceMgr &srcMgr; 1154 Operator &op; 1155 StringRef curBuffer; 1156 const char *curPtr; 1157 }; 1158 } // end anonymous namespace 1159 1160 FormatLexer::FormatLexer(llvm::SourceMgr &mgr, Operator &op) 1161 : srcMgr(mgr), op(op) { 1162 curBuffer = srcMgr.getMemoryBuffer(mgr.getMainFileID())->getBuffer(); 1163 curPtr = curBuffer.begin(); 1164 } 1165 1166 Token FormatLexer::emitError(llvm::SMLoc loc, const Twine &msg) { 1167 srcMgr.PrintMessage(loc, llvm::SourceMgr::DK_Error, msg); 1168 llvm::SrcMgr.PrintMessage(op.getLoc()[0], llvm::SourceMgr::DK_Note, 1169 "in custom assembly format for this operation"); 1170 return formToken(Token::error, loc.getPointer()); 1171 } 1172 Token FormatLexer::emitErrorAndNote(llvm::SMLoc loc, const Twine &msg, 1173 const Twine ¬e) { 1174 srcMgr.PrintMessage(loc, llvm::SourceMgr::DK_Error, msg); 1175 llvm::SrcMgr.PrintMessage(op.getLoc()[0], llvm::SourceMgr::DK_Note, 1176 "in custom assembly format for this operation"); 1177 srcMgr.PrintMessage(loc, llvm::SourceMgr::DK_Note, note); 1178 return formToken(Token::error, loc.getPointer()); 1179 } 1180 Token FormatLexer::emitError(const char *loc, const Twine &msg) { 1181 return emitError(llvm::SMLoc::getFromPointer(loc), msg); 1182 } 1183 1184 int FormatLexer::getNextChar() { 1185 char curChar = *curPtr++; 1186 switch (curChar) { 1187 default: 1188 return (unsigned char)curChar; 1189 case 0: { 1190 // A nul character in the stream is either the end of the current buffer or 1191 // a random nul in the file. Disambiguate that here. 1192 if (curPtr - 1 != curBuffer.end()) 1193 return 0; 1194 1195 // Otherwise, return end of file. 1196 --curPtr; 1197 return EOF; 1198 } 1199 case '\n': 1200 case '\r': 1201 // Handle the newline character by ignoring it and incrementing the line 1202 // count. However, be careful about 'dos style' files with \n\r in them. 1203 // Only treat a \n\r or \r\n as a single line. 1204 if ((*curPtr == '\n' || (*curPtr == '\r')) && *curPtr != curChar) 1205 ++curPtr; 1206 return '\n'; 1207 } 1208 } 1209 1210 Token FormatLexer::lexToken() { 1211 const char *tokStart = curPtr; 1212 1213 // This always consumes at least one character. 1214 int curChar = getNextChar(); 1215 switch (curChar) { 1216 default: 1217 // Handle identifiers: [a-zA-Z_] 1218 if (isalpha(curChar) || curChar == '_') 1219 return lexIdentifier(tokStart); 1220 1221 // Unknown character, emit an error. 1222 return emitError(tokStart, "unexpected character"); 1223 case EOF: 1224 // Return EOF denoting the end of lexing. 1225 return formToken(Token::eof, tokStart); 1226 1227 // Lex punctuation. 1228 case '^': 1229 return formToken(Token::caret, tokStart); 1230 case ',': 1231 return formToken(Token::comma, tokStart); 1232 case '=': 1233 return formToken(Token::equal, tokStart); 1234 case '?': 1235 return formToken(Token::question, tokStart); 1236 case '(': 1237 return formToken(Token::l_paren, tokStart); 1238 case ')': 1239 return formToken(Token::r_paren, tokStart); 1240 1241 // Ignore whitespace characters. 1242 case 0: 1243 case ' ': 1244 case '\t': 1245 case '\n': 1246 return lexToken(); 1247 1248 case '`': 1249 return lexLiteral(tokStart); 1250 case '$': 1251 return lexVariable(tokStart); 1252 } 1253 } 1254 1255 Token FormatLexer::lexLiteral(const char *tokStart) { 1256 assert(curPtr[-1] == '`'); 1257 1258 // Lex a literal surrounded by ``. 1259 while (const char curChar = *curPtr++) { 1260 if (curChar == '`') 1261 return formToken(Token::literal, tokStart); 1262 } 1263 return emitError(curPtr - 1, "unexpected end of file in literal"); 1264 } 1265 1266 Token FormatLexer::lexVariable(const char *tokStart) { 1267 if (!isalpha(curPtr[0]) && curPtr[0] != '_') 1268 return emitError(curPtr - 1, "expected variable name"); 1269 1270 // Otherwise, consume the rest of the characters. 1271 while (isalnum(*curPtr) || *curPtr == '_') 1272 ++curPtr; 1273 return formToken(Token::variable, tokStart); 1274 } 1275 1276 Token FormatLexer::lexIdentifier(const char *tokStart) { 1277 // Match the rest of the identifier regex: [0-9a-zA-Z_\-]* 1278 while (isalnum(*curPtr) || *curPtr == '_' || *curPtr == '-') 1279 ++curPtr; 1280 1281 // Check to see if this identifier is a keyword. 1282 StringRef str(tokStart, curPtr - tokStart); 1283 Token::Kind kind = 1284 llvm::StringSwitch<Token::Kind>(str) 1285 .Case("attr-dict", Token::kw_attr_dict) 1286 .Case("attr-dict-with-keyword", Token::kw_attr_dict_w_keyword) 1287 .Case("functional-type", Token::kw_functional_type) 1288 .Case("operands", Token::kw_operands) 1289 .Case("results", Token::kw_results) 1290 .Case("successors", Token::kw_successors) 1291 .Case("type", Token::kw_type) 1292 .Default(Token::identifier); 1293 return Token(kind, str); 1294 } 1295 1296 //===----------------------------------------------------------------------===// 1297 // FormatParser 1298 //===----------------------------------------------------------------------===// 1299 1300 /// Function to find an element within the given range that has the same name as 1301 /// 'name'. 1302 template <typename RangeT> static auto findArg(RangeT &&range, StringRef name) { 1303 auto it = llvm::find_if(range, [=](auto &arg) { return arg.name == name; }); 1304 return it != range.end() ? &*it : nullptr; 1305 } 1306 1307 namespace { 1308 /// This class implements a parser for an instance of an operation assembly 1309 /// format. 1310 class FormatParser { 1311 public: 1312 FormatParser(llvm::SourceMgr &mgr, OperationFormat &format, Operator &op) 1313 : lexer(mgr, op), curToken(lexer.lexToken()), fmt(format), op(op), 1314 seenOperandTypes(op.getNumOperands()), 1315 seenResultTypes(op.getNumResults()) {} 1316 1317 /// Parse the operation assembly format. 1318 LogicalResult parse(); 1319 1320 private: 1321 /// This struct represents a type resolution instance. It includes a specific 1322 /// type as well as an optional transformer to apply to that type in order to 1323 /// properly resolve the type of a variable. 1324 struct TypeResolutionInstance { 1325 const NamedTypeConstraint *type; 1326 Optional<StringRef> transformer; 1327 }; 1328 1329 /// An iterator over the elements of a format group. 1330 using ElementsIterT = llvm::pointee_iterator< 1331 std::vector<std::unique_ptr<Element>>::const_iterator>; 1332 1333 /// Verify the state of operation attributes within the format. 1334 LogicalResult verifyAttributes(llvm::SMLoc loc); 1335 /// Verify the attribute elements at the back of the given stack of iterators. 1336 LogicalResult verifyAttributes( 1337 llvm::SMLoc loc, 1338 SmallVectorImpl<std::pair<ElementsIterT, ElementsIterT>> &iteratorStack); 1339 1340 /// Verify the state of operation operands within the format. 1341 LogicalResult 1342 verifyOperands(llvm::SMLoc loc, 1343 llvm::StringMap<TypeResolutionInstance> &variableTyResolver); 1344 1345 /// Verify the state of operation results within the format. 1346 LogicalResult 1347 verifyResults(llvm::SMLoc loc, 1348 llvm::StringMap<TypeResolutionInstance> &variableTyResolver); 1349 1350 /// Verify the state of operation successors within the format. 1351 LogicalResult verifySuccessors(llvm::SMLoc loc); 1352 1353 /// Given the values of an `AllTypesMatch` trait, check for inferable type 1354 /// resolution. 1355 void handleAllTypesMatchConstraint( 1356 ArrayRef<StringRef> values, 1357 llvm::StringMap<TypeResolutionInstance> &variableTyResolver); 1358 /// Check for inferable type resolution given all operands, and or results, 1359 /// have the same type. If 'includeResults' is true, the results also have the 1360 /// same type as all of the operands. 1361 void handleSameTypesConstraint( 1362 llvm::StringMap<TypeResolutionInstance> &variableTyResolver, 1363 bool includeResults); 1364 1365 /// Returns an argument with the given name that has been seen within the 1366 /// format. 1367 const NamedTypeConstraint *findSeenArg(StringRef name); 1368 1369 /// Parse a specific element. 1370 LogicalResult parseElement(std::unique_ptr<Element> &element, 1371 bool isTopLevel); 1372 LogicalResult parseVariable(std::unique_ptr<Element> &element, 1373 bool isTopLevel); 1374 LogicalResult parseDirective(std::unique_ptr<Element> &element, 1375 bool isTopLevel); 1376 LogicalResult parseLiteral(std::unique_ptr<Element> &element); 1377 LogicalResult parseOptional(std::unique_ptr<Element> &element, 1378 bool isTopLevel); 1379 LogicalResult parseOptionalChildElement( 1380 std::vector<std::unique_ptr<Element>> &childElements, 1381 SmallPtrSetImpl<const NamedTypeConstraint *> &seenVariables, 1382 Optional<unsigned> &anchorIdx); 1383 1384 /// Parse the various different directives. 1385 LogicalResult parseAttrDictDirective(std::unique_ptr<Element> &element, 1386 llvm::SMLoc loc, bool isTopLevel, 1387 bool withKeyword); 1388 LogicalResult parseFunctionalTypeDirective(std::unique_ptr<Element> &element, 1389 Token tok, bool isTopLevel); 1390 LogicalResult parseOperandsDirective(std::unique_ptr<Element> &element, 1391 llvm::SMLoc loc, bool isTopLevel); 1392 LogicalResult parseResultsDirective(std::unique_ptr<Element> &element, 1393 llvm::SMLoc loc, bool isTopLevel); 1394 LogicalResult parseSuccessorsDirective(std::unique_ptr<Element> &element, 1395 llvm::SMLoc loc, bool isTopLevel); 1396 LogicalResult parseTypeDirective(std::unique_ptr<Element> &element, Token tok, 1397 bool isTopLevel); 1398 LogicalResult parseTypeDirectiveOperand(std::unique_ptr<Element> &element); 1399 1400 //===--------------------------------------------------------------------===// 1401 // Lexer Utilities 1402 //===--------------------------------------------------------------------===// 1403 1404 /// Advance the current lexer onto the next token. 1405 void consumeToken() { 1406 assert(curToken.getKind() != Token::eof && 1407 curToken.getKind() != Token::error && 1408 "shouldn't advance past EOF or errors"); 1409 curToken = lexer.lexToken(); 1410 } 1411 LogicalResult parseToken(Token::Kind kind, const Twine &msg) { 1412 if (curToken.getKind() != kind) 1413 return emitError(curToken.getLoc(), msg); 1414 consumeToken(); 1415 return success(); 1416 } 1417 LogicalResult emitError(llvm::SMLoc loc, const Twine &msg) { 1418 lexer.emitError(loc, msg); 1419 return failure(); 1420 } 1421 LogicalResult emitErrorAndNote(llvm::SMLoc loc, const Twine &msg, 1422 const Twine ¬e) { 1423 lexer.emitErrorAndNote(loc, msg, note); 1424 return failure(); 1425 } 1426 1427 //===--------------------------------------------------------------------===// 1428 // Fields 1429 //===--------------------------------------------------------------------===// 1430 1431 FormatLexer lexer; 1432 Token curToken; 1433 OperationFormat &fmt; 1434 Operator &op; 1435 1436 // The following are various bits of format state used for verification 1437 // during parsing. 1438 bool hasAllOperands = false, hasAttrDict = false; 1439 bool hasAllSuccessors = false; 1440 llvm::SmallBitVector seenOperandTypes, seenResultTypes; 1441 llvm::DenseSet<const NamedTypeConstraint *> seenOperands; 1442 llvm::DenseSet<const NamedAttribute *> seenAttrs; 1443 llvm::DenseSet<const NamedSuccessor *> seenSuccessors; 1444 llvm::DenseSet<const NamedTypeConstraint *> optionalVariables; 1445 }; 1446 } // end anonymous namespace 1447 1448 LogicalResult FormatParser::parse() { 1449 llvm::SMLoc loc = curToken.getLoc(); 1450 1451 // Parse each of the format elements into the main format. 1452 while (curToken.getKind() != Token::eof) { 1453 std::unique_ptr<Element> element; 1454 if (failed(parseElement(element, /*isTopLevel=*/true))) 1455 return failure(); 1456 fmt.elements.push_back(std::move(element)); 1457 } 1458 1459 // Check that the attribute dictionary is in the format. 1460 if (!hasAttrDict) 1461 return emitError(loc, "'attr-dict' directive not found in " 1462 "custom assembly format"); 1463 1464 // Check for any type traits that we can use for inferring types. 1465 llvm::StringMap<TypeResolutionInstance> variableTyResolver; 1466 for (const OpTrait &trait : op.getTraits()) { 1467 const llvm::Record &def = trait.getDef(); 1468 if (def.isSubClassOf("AllTypesMatch")) { 1469 handleAllTypesMatchConstraint(def.getValueAsListOfStrings("values"), 1470 variableTyResolver); 1471 } else if (def.getName() == "SameTypeOperands") { 1472 handleSameTypesConstraint(variableTyResolver, /*includeResults=*/false); 1473 } else if (def.getName() == "SameOperandsAndResultType") { 1474 handleSameTypesConstraint(variableTyResolver, /*includeResults=*/true); 1475 } else if (def.isSubClassOf("TypesMatchWith")) { 1476 if (const auto *lhsArg = findSeenArg(def.getValueAsString("lhs"))) 1477 variableTyResolver[def.getValueAsString("rhs")] = { 1478 lhsArg, def.getValueAsString("transformer")}; 1479 } 1480 } 1481 1482 // Verify the state of the various operation components. 1483 if (failed(verifyAttributes(loc)) || 1484 failed(verifyResults(loc, variableTyResolver)) || 1485 failed(verifyOperands(loc, variableTyResolver)) || 1486 failed(verifySuccessors(loc))) 1487 return failure(); 1488 1489 // Check to see if we are formatting all of the operands. 1490 fmt.allOperands = llvm::any_of(fmt.elements, [](auto &elt) { 1491 return isa<OperandsDirective>(elt.get()); 1492 }); 1493 return success(); 1494 } 1495 1496 LogicalResult FormatParser::verifyAttributes(llvm::SMLoc loc) { 1497 // Check that there are no `:` literals after an attribute without a constant 1498 // type. The attribute grammar contains an optional trailing colon type, which 1499 // can lead to unexpected and generally unintended behavior. Given that, it is 1500 // better to just error out here instead. 1501 using ElementsIterT = llvm::pointee_iterator< 1502 std::vector<std::unique_ptr<Element>>::const_iterator>; 1503 SmallVector<std::pair<ElementsIterT, ElementsIterT>, 1> iteratorStack; 1504 iteratorStack.emplace_back(fmt.elements.begin(), fmt.elements.end()); 1505 while (!iteratorStack.empty()) 1506 if (failed(verifyAttributes(loc, iteratorStack))) 1507 return failure(); 1508 return success(); 1509 } 1510 /// Verify the attribute elements at the back of the given stack of iterators. 1511 LogicalResult FormatParser::verifyAttributes( 1512 llvm::SMLoc loc, 1513 SmallVectorImpl<std::pair<ElementsIterT, ElementsIterT>> &iteratorStack) { 1514 auto &stackIt = iteratorStack.back(); 1515 ElementsIterT &it = stackIt.first, e = stackIt.second; 1516 while (it != e) { 1517 Element *element = &*(it++); 1518 1519 // Traverse into optional groups. 1520 if (auto *optional = dyn_cast<OptionalElement>(element)) { 1521 auto elements = optional->getElements(); 1522 iteratorStack.emplace_back(elements.begin(), elements.end()); 1523 return success(); 1524 } 1525 1526 // We are checking for an attribute element followed by a `:`, so there is 1527 // no need to check the end. 1528 if (it == e && iteratorStack.size() == 1) 1529 break; 1530 1531 // Check for an attribute with a constant type builder, followed by a `:`. 1532 auto *prevAttr = dyn_cast<AttributeVariable>(element); 1533 if (!prevAttr || prevAttr->getTypeBuilder()) 1534 continue; 1535 1536 // Check the next iterator within the stack for literal elements. 1537 for (auto &nextItPair : iteratorStack) { 1538 ElementsIterT nextIt = nextItPair.first, nextE = nextItPair.second; 1539 for (; nextIt != nextE; ++nextIt) { 1540 // Skip any trailing optional groups or attribute dictionaries. 1541 if (isa<AttrDictDirective>(*nextIt) || isa<OptionalElement>(*nextIt)) 1542 continue; 1543 1544 // We are only interested in `:` literals. 1545 auto *literal = dyn_cast<LiteralElement>(&*nextIt); 1546 if (!literal || literal->getLiteral() != ":") 1547 break; 1548 1549 // TODO: Use the location of the literal element itself. 1550 return emitError( 1551 loc, llvm::formatv("format ambiguity caused by `:` literal found " 1552 "after attribute `{0}` which does not have " 1553 "a buildable type", 1554 prevAttr->getVar()->name)); 1555 } 1556 } 1557 } 1558 iteratorStack.pop_back(); 1559 return success(); 1560 } 1561 1562 LogicalResult FormatParser::verifyOperands( 1563 llvm::SMLoc loc, 1564 llvm::StringMap<TypeResolutionInstance> &variableTyResolver) { 1565 // Check that all of the operands are within the format, and their types can 1566 // be inferred. 1567 auto &buildableTypes = fmt.buildableTypes; 1568 for (unsigned i = 0, e = op.getNumOperands(); i != e; ++i) { 1569 NamedTypeConstraint &operand = op.getOperand(i); 1570 1571 // Check that the operand itself is in the format. 1572 if (!hasAllOperands && !seenOperands.count(&operand)) { 1573 return emitErrorAndNote(loc, 1574 "operand #" + Twine(i) + ", named '" + 1575 operand.name + "', not found", 1576 "suggest adding a '$" + operand.name + 1577 "' directive to the custom assembly format"); 1578 } 1579 1580 // Check that the operand type is in the format, or that it can be inferred. 1581 if (fmt.allOperandTypes || seenOperandTypes.test(i)) 1582 continue; 1583 1584 // Check to see if we can infer this type from another variable. 1585 auto varResolverIt = variableTyResolver.find(op.getOperand(i).name); 1586 if (varResolverIt != variableTyResolver.end()) { 1587 fmt.operandTypes[i].setVariable(varResolverIt->second.type, 1588 varResolverIt->second.transformer); 1589 continue; 1590 } 1591 1592 // Similarly to results, allow a custom builder for resolving the type if 1593 // we aren't using the 'operands' directive. 1594 Optional<StringRef> builder = operand.constraint.getBuilderCall(); 1595 if (!builder || (hasAllOperands && operand.isVariableLength())) { 1596 return emitErrorAndNote( 1597 loc, 1598 "type of operand #" + Twine(i) + ", named '" + operand.name + 1599 "', is not buildable and a buildable type cannot be inferred", 1600 "suggest adding a type constraint to the operation or adding a " 1601 "'type($" + 1602 operand.name + ")' directive to the " + "custom assembly format"); 1603 } 1604 auto it = buildableTypes.insert({*builder, buildableTypes.size()}); 1605 fmt.operandTypes[i].setBuilderIdx(it.first->second); 1606 } 1607 return success(); 1608 } 1609 1610 LogicalResult FormatParser::verifyResults( 1611 llvm::SMLoc loc, 1612 llvm::StringMap<TypeResolutionInstance> &variableTyResolver) { 1613 // If we format all of the types together, there is nothing to check. 1614 if (fmt.allResultTypes) 1615 return success(); 1616 1617 // Check that all of the result types can be inferred. 1618 auto &buildableTypes = fmt.buildableTypes; 1619 for (unsigned i = 0, e = op.getNumResults(); i != e; ++i) { 1620 if (seenResultTypes.test(i)) 1621 continue; 1622 1623 // Check to see if we can infer this type from another variable. 1624 auto varResolverIt = variableTyResolver.find(op.getResultName(i)); 1625 if (varResolverIt != variableTyResolver.end()) { 1626 fmt.resultTypes[i].setVariable(varResolverIt->second.type, 1627 varResolverIt->second.transformer); 1628 continue; 1629 } 1630 1631 // If the result is not variable length, allow for the case where the type 1632 // has a builder that we can use. 1633 NamedTypeConstraint &result = op.getResult(i); 1634 Optional<StringRef> builder = result.constraint.getBuilderCall(); 1635 if (!builder || result.isVariableLength()) { 1636 return emitErrorAndNote( 1637 loc, 1638 "type of result #" + Twine(i) + ", named '" + result.name + 1639 "', is not buildable and a buildable type cannot be inferred", 1640 "suggest adding a type constraint to the operation or adding a " 1641 "'type($" + 1642 result.name + ")' directive to the " + "custom assembly format"); 1643 } 1644 // Note in the format that this result uses the custom builder. 1645 auto it = buildableTypes.insert({*builder, buildableTypes.size()}); 1646 fmt.resultTypes[i].setBuilderIdx(it.first->second); 1647 } 1648 return success(); 1649 } 1650 1651 LogicalResult FormatParser::verifySuccessors(llvm::SMLoc loc) { 1652 // Check that all of the successors are within the format. 1653 if (hasAllSuccessors) 1654 return success(); 1655 1656 for (unsigned i = 0, e = op.getNumSuccessors(); i != e; ++i) { 1657 const NamedSuccessor &successor = op.getSuccessor(i); 1658 if (!seenSuccessors.count(&successor)) { 1659 return emitErrorAndNote(loc, 1660 "successor #" + Twine(i) + ", named '" + 1661 successor.name + "', not found", 1662 "suggest adding a '$" + successor.name + 1663 "' directive to the custom assembly format"); 1664 } 1665 } 1666 return success(); 1667 } 1668 1669 void FormatParser::handleAllTypesMatchConstraint( 1670 ArrayRef<StringRef> values, 1671 llvm::StringMap<TypeResolutionInstance> &variableTyResolver) { 1672 for (unsigned i = 0, e = values.size(); i != e; ++i) { 1673 // Check to see if this value matches a resolved operand or result type. 1674 const NamedTypeConstraint *arg = findSeenArg(values[i]); 1675 if (!arg) 1676 continue; 1677 1678 // Mark this value as the type resolver for the other variables. 1679 for (unsigned j = 0; j != i; ++j) 1680 variableTyResolver[values[j]] = {arg, llvm::None}; 1681 for (unsigned j = i + 1; j != e; ++j) 1682 variableTyResolver[values[j]] = {arg, llvm::None}; 1683 } 1684 } 1685 1686 void FormatParser::handleSameTypesConstraint( 1687 llvm::StringMap<TypeResolutionInstance> &variableTyResolver, 1688 bool includeResults) { 1689 const NamedTypeConstraint *resolver = nullptr; 1690 int resolvedIt = -1; 1691 1692 // Check to see if there is an operand or result to use for the resolution. 1693 if ((resolvedIt = seenOperandTypes.find_first()) != -1) 1694 resolver = &op.getOperand(resolvedIt); 1695 else if (includeResults && (resolvedIt = seenResultTypes.find_first()) != -1) 1696 resolver = &op.getResult(resolvedIt); 1697 else 1698 return; 1699 1700 // Set the resolvers for each operand and result. 1701 for (unsigned i = 0, e = op.getNumOperands(); i != e; ++i) 1702 if (!seenOperandTypes.test(i) && !op.getOperand(i).name.empty()) 1703 variableTyResolver[op.getOperand(i).name] = {resolver, llvm::None}; 1704 if (includeResults) { 1705 for (unsigned i = 0, e = op.getNumResults(); i != e; ++i) 1706 if (!seenResultTypes.test(i) && !op.getResultName(i).empty()) 1707 variableTyResolver[op.getResultName(i)] = {resolver, llvm::None}; 1708 } 1709 } 1710 1711 const NamedTypeConstraint *FormatParser::findSeenArg(StringRef name) { 1712 if (auto *arg = findArg(op.getOperands(), name)) 1713 return seenOperandTypes.test(arg - op.operand_begin()) ? arg : nullptr; 1714 if (auto *arg = findArg(op.getResults(), name)) 1715 return seenResultTypes.test(arg - op.result_begin()) ? arg : nullptr; 1716 return nullptr; 1717 } 1718 1719 LogicalResult FormatParser::parseElement(std::unique_ptr<Element> &element, 1720 bool isTopLevel) { 1721 // Directives. 1722 if (curToken.isKeyword()) 1723 return parseDirective(element, isTopLevel); 1724 // Literals. 1725 if (curToken.getKind() == Token::literal) 1726 return parseLiteral(element); 1727 // Optionals. 1728 if (curToken.getKind() == Token::l_paren) 1729 return parseOptional(element, isTopLevel); 1730 // Variables. 1731 if (curToken.getKind() == Token::variable) 1732 return parseVariable(element, isTopLevel); 1733 return emitError(curToken.getLoc(), 1734 "expected directive, literal, variable, or optional group"); 1735 } 1736 1737 LogicalResult FormatParser::parseVariable(std::unique_ptr<Element> &element, 1738 bool isTopLevel) { 1739 Token varTok = curToken; 1740 consumeToken(); 1741 1742 StringRef name = varTok.getSpelling().drop_front(); 1743 llvm::SMLoc loc = varTok.getLoc(); 1744 1745 // Check that the parsed argument is something actually registered on the 1746 // op. 1747 /// Attributes 1748 if (const NamedAttribute *attr = findArg(op.getAttributes(), name)) { 1749 if (isTopLevel && !seenAttrs.insert(attr).second) 1750 return emitError(loc, "attribute '" + name + "' is already bound"); 1751 element = std::make_unique<AttributeVariable>(attr); 1752 return success(); 1753 } 1754 /// Operands 1755 if (const NamedTypeConstraint *operand = findArg(op.getOperands(), name)) { 1756 if (isTopLevel) { 1757 if (hasAllOperands || !seenOperands.insert(operand).second) 1758 return emitError(loc, "operand '" + name + "' is already bound"); 1759 } 1760 element = std::make_unique<OperandVariable>(operand); 1761 return success(); 1762 } 1763 /// Results. 1764 if (const auto *result = findArg(op.getResults(), name)) { 1765 if (isTopLevel) 1766 return emitError(loc, "results can not be used at the top level"); 1767 element = std::make_unique<ResultVariable>(result); 1768 return success(); 1769 } 1770 /// Successors. 1771 if (const auto *successor = findArg(op.getSuccessors(), name)) { 1772 if (!isTopLevel) 1773 return emitError(loc, "successors can only be used at the top level"); 1774 if (hasAllSuccessors || !seenSuccessors.insert(successor).second) 1775 return emitError(loc, "successor '" + name + "' is already bound"); 1776 element = std::make_unique<SuccessorVariable>(successor); 1777 return success(); 1778 } 1779 return emitError( 1780 loc, "expected variable to refer to an argument, result, or successor"); 1781 } 1782 1783 LogicalResult FormatParser::parseDirective(std::unique_ptr<Element> &element, 1784 bool isTopLevel) { 1785 Token dirTok = curToken; 1786 consumeToken(); 1787 1788 switch (dirTok.getKind()) { 1789 case Token::kw_attr_dict: 1790 return parseAttrDictDirective(element, dirTok.getLoc(), isTopLevel, 1791 /*withKeyword=*/false); 1792 case Token::kw_attr_dict_w_keyword: 1793 return parseAttrDictDirective(element, dirTok.getLoc(), isTopLevel, 1794 /*withKeyword=*/true); 1795 case Token::kw_functional_type: 1796 return parseFunctionalTypeDirective(element, dirTok, isTopLevel); 1797 case Token::kw_operands: 1798 return parseOperandsDirective(element, dirTok.getLoc(), isTopLevel); 1799 case Token::kw_results: 1800 return parseResultsDirective(element, dirTok.getLoc(), isTopLevel); 1801 case Token::kw_successors: 1802 return parseSuccessorsDirective(element, dirTok.getLoc(), isTopLevel); 1803 case Token::kw_type: 1804 return parseTypeDirective(element, dirTok, isTopLevel); 1805 1806 default: 1807 llvm_unreachable("unknown directive token"); 1808 } 1809 } 1810 1811 LogicalResult FormatParser::parseLiteral(std::unique_ptr<Element> &element) { 1812 Token literalTok = curToken; 1813 consumeToken(); 1814 1815 // Check that the parsed literal is valid. 1816 StringRef value = literalTok.getSpelling().drop_front().drop_back(); 1817 if (!LiteralElement::isValidLiteral(value)) 1818 return emitError(literalTok.getLoc(), "expected valid literal"); 1819 1820 element = std::make_unique<LiteralElement>(value); 1821 return success(); 1822 } 1823 1824 LogicalResult FormatParser::parseOptional(std::unique_ptr<Element> &element, 1825 bool isTopLevel) { 1826 llvm::SMLoc curLoc = curToken.getLoc(); 1827 if (!isTopLevel) 1828 return emitError(curLoc, "optional groups can only be used as top-level " 1829 "elements"); 1830 consumeToken(); 1831 1832 // Parse the child elements for this optional group. 1833 std::vector<std::unique_ptr<Element>> elements; 1834 SmallPtrSet<const NamedTypeConstraint *, 8> seenVariables; 1835 Optional<unsigned> anchorIdx; 1836 do { 1837 if (failed(parseOptionalChildElement(elements, seenVariables, anchorIdx))) 1838 return failure(); 1839 } while (curToken.getKind() != Token::r_paren); 1840 consumeToken(); 1841 if (failed(parseToken(Token::question, "expected '?' after optional group"))) 1842 return failure(); 1843 1844 // The optional group is required to have an anchor. 1845 if (!anchorIdx) 1846 return emitError(curLoc, "optional group specified no anchor element"); 1847 1848 // The first element of the group must be one that can be parsed/printed in an 1849 // optional fashion. 1850 if (!isa<LiteralElement>(&*elements.front()) && 1851 !isa<OperandVariable>(&*elements.front())) 1852 return emitError(curLoc, "first element of an operand group must be a " 1853 "literal or operand"); 1854 1855 // After parsing all of the elements, ensure that all type directives refer 1856 // only to elements within the group. 1857 auto checkTypeOperand = [&](Element *typeEle) { 1858 auto *opVar = dyn_cast<OperandVariable>(typeEle); 1859 const NamedTypeConstraint *var = opVar ? opVar->getVar() : nullptr; 1860 if (!seenVariables.count(var)) 1861 return emitError(curLoc, "type directive can only refer to variables " 1862 "within the optional group"); 1863 return success(); 1864 }; 1865 for (auto &ele : elements) { 1866 if (auto *typeEle = dyn_cast<TypeDirective>(ele.get())) { 1867 if (failed(checkTypeOperand(typeEle->getOperand()))) 1868 return failure(); 1869 } else if (auto *typeEle = dyn_cast<FunctionalTypeDirective>(ele.get())) { 1870 if (failed(checkTypeOperand(typeEle->getInputs())) || 1871 failed(checkTypeOperand(typeEle->getResults()))) 1872 return failure(); 1873 } 1874 } 1875 1876 optionalVariables.insert(seenVariables.begin(), seenVariables.end()); 1877 element = std::make_unique<OptionalElement>(std::move(elements), *anchorIdx); 1878 return success(); 1879 } 1880 1881 LogicalResult FormatParser::parseOptionalChildElement( 1882 std::vector<std::unique_ptr<Element>> &childElements, 1883 SmallPtrSetImpl<const NamedTypeConstraint *> &seenVariables, 1884 Optional<unsigned> &anchorIdx) { 1885 llvm::SMLoc childLoc = curToken.getLoc(); 1886 childElements.push_back({}); 1887 if (failed(parseElement(childElements.back(), /*isTopLevel=*/true))) 1888 return failure(); 1889 1890 // Check to see if this element is the anchor of the optional group. 1891 bool isAnchor = curToken.getKind() == Token::caret; 1892 if (isAnchor) { 1893 if (anchorIdx) 1894 return emitError(childLoc, "only one element can be marked as the anchor " 1895 "of an optional group"); 1896 anchorIdx = childElements.size() - 1; 1897 consumeToken(); 1898 } 1899 1900 return TypeSwitch<Element *, LogicalResult>(childElements.back().get()) 1901 // All attributes can be within the optional group, but only optional 1902 // attributes can be the anchor. 1903 .Case([&](AttributeVariable *attrEle) { 1904 if (isAnchor && !attrEle->getVar()->attr.isOptional()) 1905 return emitError(childLoc, "only optional attributes can be used to " 1906 "anchor an optional group"); 1907 return success(); 1908 }) 1909 // Only optional-like(i.e. variadic) operands can be within an optional 1910 // group. 1911 .Case<OperandVariable>([&](OperandVariable *ele) { 1912 if (!ele->getVar()->isVariableLength()) 1913 return emitError(childLoc, "only variable length operands can be " 1914 "used within an optional group"); 1915 seenVariables.insert(ele->getVar()); 1916 return success(); 1917 }) 1918 // Literals and type directives may be used, but they can't anchor the 1919 // group. 1920 .Case<LiteralElement, TypeDirective, FunctionalTypeDirective>( 1921 [&](Element *) { 1922 if (isAnchor) 1923 return emitError(childLoc, "only variables can be used to anchor " 1924 "an optional group"); 1925 return success(); 1926 }) 1927 .Default([&](Element *) { 1928 return emitError(childLoc, "only literals, types, and variables can be " 1929 "used within an optional group"); 1930 }); 1931 } 1932 1933 LogicalResult 1934 FormatParser::parseAttrDictDirective(std::unique_ptr<Element> &element, 1935 llvm::SMLoc loc, bool isTopLevel, 1936 bool withKeyword) { 1937 if (!isTopLevel) 1938 return emitError(loc, "'attr-dict' directive can only be used as a " 1939 "top-level directive"); 1940 if (hasAttrDict) 1941 return emitError(loc, "'attr-dict' directive has already been seen"); 1942 1943 hasAttrDict = true; 1944 element = std::make_unique<AttrDictDirective>(withKeyword); 1945 return success(); 1946 } 1947 1948 LogicalResult 1949 FormatParser::parseFunctionalTypeDirective(std::unique_ptr<Element> &element, 1950 Token tok, bool isTopLevel) { 1951 llvm::SMLoc loc = tok.getLoc(); 1952 if (!isTopLevel) 1953 return emitError( 1954 loc, "'functional-type' is only valid as a top-level directive"); 1955 1956 // Parse the main operand. 1957 std::unique_ptr<Element> inputs, results; 1958 if (failed(parseToken(Token::l_paren, "expected '(' before argument list")) || 1959 failed(parseTypeDirectiveOperand(inputs)) || 1960 failed(parseToken(Token::comma, "expected ',' after inputs argument")) || 1961 failed(parseTypeDirectiveOperand(results)) || 1962 failed(parseToken(Token::r_paren, "expected ')' after argument list"))) 1963 return failure(); 1964 element = std::make_unique<FunctionalTypeDirective>(std::move(inputs), 1965 std::move(results)); 1966 return success(); 1967 } 1968 1969 LogicalResult 1970 FormatParser::parseOperandsDirective(std::unique_ptr<Element> &element, 1971 llvm::SMLoc loc, bool isTopLevel) { 1972 if (isTopLevel && (hasAllOperands || !seenOperands.empty())) 1973 return emitError(loc, "'operands' directive creates overlap in format"); 1974 hasAllOperands = true; 1975 element = std::make_unique<OperandsDirective>(); 1976 return success(); 1977 } 1978 1979 LogicalResult 1980 FormatParser::parseResultsDirective(std::unique_ptr<Element> &element, 1981 llvm::SMLoc loc, bool isTopLevel) { 1982 if (isTopLevel) 1983 return emitError(loc, "'results' directive can not be used as a " 1984 "top-level directive"); 1985 element = std::make_unique<ResultsDirective>(); 1986 return success(); 1987 } 1988 1989 LogicalResult 1990 FormatParser::parseSuccessorsDirective(std::unique_ptr<Element> &element, 1991 llvm::SMLoc loc, bool isTopLevel) { 1992 if (!isTopLevel) 1993 return emitError(loc, 1994 "'successors' is only valid as a top-level directive"); 1995 if (hasAllSuccessors || !seenSuccessors.empty()) 1996 return emitError(loc, "'successors' directive creates overlap in format"); 1997 hasAllSuccessors = true; 1998 element = std::make_unique<SuccessorsDirective>(); 1999 return success(); 2000 } 2001 2002 LogicalResult 2003 FormatParser::parseTypeDirective(std::unique_ptr<Element> &element, Token tok, 2004 bool isTopLevel) { 2005 llvm::SMLoc loc = tok.getLoc(); 2006 if (!isTopLevel) 2007 return emitError(loc, "'type' is only valid as a top-level directive"); 2008 2009 std::unique_ptr<Element> operand; 2010 if (failed(parseToken(Token::l_paren, "expected '(' before argument list")) || 2011 failed(parseTypeDirectiveOperand(operand)) || 2012 failed(parseToken(Token::r_paren, "expected ')' after argument list"))) 2013 return failure(); 2014 element = std::make_unique<TypeDirective>(std::move(operand)); 2015 return success(); 2016 } 2017 2018 LogicalResult 2019 FormatParser::parseTypeDirectiveOperand(std::unique_ptr<Element> &element) { 2020 llvm::SMLoc loc = curToken.getLoc(); 2021 if (failed(parseElement(element, /*isTopLevel=*/false))) 2022 return failure(); 2023 if (isa<LiteralElement>(element.get())) 2024 return emitError( 2025 loc, "'type' directive operand expects variable or directive operand"); 2026 2027 if (auto *var = dyn_cast<OperandVariable>(element.get())) { 2028 unsigned opIdx = var->getVar() - op.operand_begin(); 2029 if (fmt.allOperandTypes || seenOperandTypes.test(opIdx)) 2030 return emitError(loc, "'type' of '" + var->getVar()->name + 2031 "' is already bound"); 2032 seenOperandTypes.set(opIdx); 2033 } else if (auto *var = dyn_cast<ResultVariable>(element.get())) { 2034 unsigned resIdx = var->getVar() - op.result_begin(); 2035 if (fmt.allResultTypes || seenResultTypes.test(resIdx)) 2036 return emitError(loc, "'type' of '" + var->getVar()->name + 2037 "' is already bound"); 2038 seenResultTypes.set(resIdx); 2039 } else if (isa<OperandsDirective>(&*element)) { 2040 if (fmt.allOperandTypes || seenOperandTypes.any()) 2041 return emitError(loc, "'operands' 'type' is already bound"); 2042 fmt.allOperandTypes = true; 2043 } else if (isa<ResultsDirective>(&*element)) { 2044 if (fmt.allResultTypes || seenResultTypes.any()) 2045 return emitError(loc, "'results' 'type' is already bound"); 2046 fmt.allResultTypes = true; 2047 } else { 2048 return emitError(loc, "invalid argument to 'type' directive"); 2049 } 2050 return success(); 2051 } 2052 2053 //===----------------------------------------------------------------------===// 2054 // Interface 2055 //===----------------------------------------------------------------------===// 2056 2057 void mlir::tblgen::generateOpFormat(const Operator &constOp, OpClass &opClass) { 2058 // TODO(riverriddle) Operator doesn't expose all necessary functionality via 2059 // the const interface. 2060 Operator &op = const_cast<Operator &>(constOp); 2061 if (!op.hasAssemblyFormat()) 2062 return; 2063 2064 // Parse the format description. 2065 llvm::SourceMgr mgr; 2066 mgr.AddNewSourceBuffer( 2067 llvm::MemoryBuffer::getMemBuffer(op.getAssemblyFormat()), llvm::SMLoc()); 2068 OperationFormat format(op); 2069 if (failed(FormatParser(mgr, format, op).parse())) { 2070 // Exit the process if format errors are treated as fatal. 2071 if (formatErrorIsFatal) { 2072 // Invoke the interrupt handlers to run the file cleanup handlers. 2073 llvm::sys::RunInterruptHandlers(); 2074 std::exit(1); 2075 } 2076 return; 2077 } 2078 2079 // Generate the printer and parser based on the parsed format. 2080 format.genParser(op, opClass); 2081 format.genPrinter(op, opClass); 2082 } 2083