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 body << " return success();\n"; 710 } 711 712 void OperationFormat::genParserTypeResolution(Operator &op, 713 OpMethodBody &body) { 714 // If any of type resolutions use transformed variables, make sure that the 715 // types of those variables are resolved. 716 SmallPtrSet<const NamedTypeConstraint *, 8> verifiedVariables; 717 FmtContext verifierFCtx; 718 for (TypeResolution &resolver : 719 llvm::concat<TypeResolution>(resultTypes, operandTypes)) { 720 Optional<StringRef> transformer = resolver.getVarTransformer(); 721 if (!transformer) 722 continue; 723 // Ensure that we don't verify the same variables twice. 724 const NamedTypeConstraint *variable = resolver.getVariable(); 725 if (!verifiedVariables.insert(variable).second) 726 continue; 727 728 auto constraint = variable->constraint; 729 body << " for (Type type : " << variable->name << "Types) {\n" 730 << " (void)type;\n" 731 << " if (!(" 732 << tgfmt(constraint.getConditionTemplate(), 733 &verifierFCtx.withSelf("type")) 734 << ")) {\n" 735 << formatv(" return parser.emitError(parser.getNameLoc()) << " 736 "\"'{0}' must be {1}, but got \" << type;\n", 737 variable->name, constraint.getDescription()) 738 << " }\n" 739 << " }\n"; 740 } 741 742 // Initialize the set of buildable types. 743 if (!buildableTypes.empty()) { 744 body << " Builder &builder = parser.getBuilder();\n"; 745 746 FmtContext typeBuilderCtx; 747 typeBuilderCtx.withBuilder("builder"); 748 for (auto &it : buildableTypes) 749 body << " Type odsBuildableType" << it.second << " = " 750 << tgfmt(it.first, &typeBuilderCtx) << ";\n"; 751 } 752 753 // Emit the code necessary for a type resolver. 754 auto emitTypeResolver = [&](TypeResolution &resolver, StringRef curVar) { 755 if (Optional<int> val = resolver.getBuilderIdx()) { 756 body << "odsBuildableType" << *val; 757 } else if (const NamedTypeConstraint *var = resolver.getVariable()) { 758 if (Optional<StringRef> tform = resolver.getVarTransformer()) 759 body << tgfmt(*tform, &FmtContext().withSelf(var->name + "Types[0]")); 760 else 761 body << var->name << "Types"; 762 } else { 763 body << curVar << "Types"; 764 } 765 }; 766 767 // Resolve each of the result types. 768 if (allResultTypes) { 769 body << " result.addTypes(allResultTypes);\n"; 770 } else { 771 for (unsigned i = 0, e = op.getNumResults(); i != e; ++i) { 772 body << " result.addTypes("; 773 emitTypeResolver(resultTypes[i], op.getResultName(i)); 774 body << ");\n"; 775 } 776 } 777 778 // Early exit if there are no operands. 779 if (op.getNumOperands() == 0) 780 return; 781 782 // Handle the case where all operand types are in one group. 783 if (allOperandTypes) { 784 // If we have all operands together, use the full operand list directly. 785 if (allOperands) { 786 body << " if (parser.resolveOperands(allOperands, allOperandTypes, " 787 "allOperandLoc, result.operands))\n" 788 " return failure();\n"; 789 return; 790 } 791 792 // Otherwise, use llvm::concat to merge the disjoint operand lists together. 793 // llvm::concat does not allow the case of a single range, so guard it here. 794 body << " if (parser.resolveOperands("; 795 if (op.getNumOperands() > 1) { 796 body << "llvm::concat<const OpAsmParser::OperandType>("; 797 llvm::interleaveComma(op.getOperands(), body, [&](auto &operand) { 798 body << operand.name << "Operands"; 799 }); 800 body << ")"; 801 } else { 802 body << op.operand_begin()->name << "Operands"; 803 } 804 body << ", allOperandTypes, parser.getNameLoc(), result.operands))\n" 805 << " return failure();\n"; 806 return; 807 } 808 // Handle the case where all of the operands were grouped together. 809 if (allOperands) { 810 body << " if (parser.resolveOperands(allOperands, "; 811 812 // Group all of the operand types together to perform the resolution all at 813 // once. Use llvm::concat to perform the merge. llvm::concat does not allow 814 // the case of a single range, so guard it here. 815 if (op.getNumOperands() > 1) { 816 body << "llvm::concat<const Type>("; 817 llvm::interleaveComma( 818 llvm::seq<int>(0, op.getNumOperands()), body, [&](int i) { 819 body << "ArrayRef<Type>("; 820 emitTypeResolver(operandTypes[i], op.getOperand(i).name); 821 body << ")"; 822 }); 823 body << ")"; 824 } else { 825 emitTypeResolver(operandTypes.front(), op.getOperand(0).name); 826 } 827 828 body << ", allOperandLoc, result.operands))\n" 829 << " return failure();\n"; 830 return; 831 } 832 833 // The final case is the one where each of the operands types are resolved 834 // separately. 835 for (unsigned i = 0, e = op.getNumOperands(); i != e; ++i) { 836 NamedTypeConstraint &operand = op.getOperand(i); 837 body << " if (parser.resolveOperands(" << operand.name << "Operands, "; 838 emitTypeResolver(operandTypes[i], operand.name); 839 840 // If this isn't a buildable type, verify the sizes match by adding the loc. 841 if (!operandTypes[i].getBuilderIdx()) 842 body << ", " << operand.name << "OperandsLoc"; 843 body << ", result.operands))\n return failure();\n"; 844 } 845 } 846 847 void OperationFormat::genParserSuccessorResolution(Operator &op, 848 OpMethodBody &body) { 849 // Check for the case where all successors were parsed. 850 bool hasAllSuccessors = llvm::any_of( 851 elements, [](auto &elt) { return isa<SuccessorsDirective>(elt.get()); }); 852 if (hasAllSuccessors) { 853 body << " result.addSuccessors(fullSuccessors);\n"; 854 return; 855 } 856 857 // Otherwise, handle each successor individually. 858 for (const NamedSuccessor &successor : op.getSuccessors()) { 859 if (successor.isVariadic()) 860 body << " result.addSuccessors(" << successor.name << "Successors);\n"; 861 else 862 body << " result.addSuccessors(" << successor.name << "Successor);\n"; 863 } 864 } 865 866 void OperationFormat::genParserVariadicSegmentResolution(Operator &op, 867 OpMethodBody &body) { 868 if (!allOperands && op.getTrait("OpTrait::AttrSizedOperandSegments")) { 869 body << " result.addAttribute(\"operand_segment_sizes\", " 870 << "builder.getI32VectorAttr({"; 871 auto interleaveFn = [&](const NamedTypeConstraint &operand) { 872 // If the operand is variadic emit the parsed size. 873 if (operand.isVariableLength()) 874 body << "static_cast<int32_t>(" << operand.name << "Operands.size())"; 875 else 876 body << "1"; 877 }; 878 llvm::interleaveComma(op.getOperands(), body, interleaveFn); 879 body << "}));\n"; 880 } 881 } 882 883 //===----------------------------------------------------------------------===// 884 // PrinterGen 885 886 /// Generate the printer for the 'attr-dict' directive. 887 static void genAttrDictPrinter(OperationFormat &fmt, Operator &op, 888 OpMethodBody &body, bool withKeyword) { 889 // Collect all of the attributes used in the format, these will be elided. 890 SmallVector<const NamedAttribute *, 1> usedAttributes; 891 for (auto &it : fmt.elements) 892 if (auto *attr = dyn_cast<AttributeVariable>(it.get())) 893 usedAttributes.push_back(attr->getVar()); 894 895 body << " p.printOptionalAttrDict" << (withKeyword ? "WithKeyword" : "") 896 << "(getAttrs(), /*elidedAttrs=*/{"; 897 // Elide the variadic segment size attributes if necessary. 898 if (!fmt.allOperands && op.getTrait("OpTrait::AttrSizedOperandSegments")) 899 body << "\"operand_segment_sizes\", "; 900 llvm::interleaveComma(usedAttributes, body, [&](const NamedAttribute *attr) { 901 body << "\"" << attr->name << "\""; 902 }); 903 body << "});\n"; 904 } 905 906 /// Generate the printer for a literal value. `shouldEmitSpace` is true if a 907 /// space should be emitted before this element. `lastWasPunctuation` is true if 908 /// the previous element was a punctuation literal. 909 static void genLiteralPrinter(StringRef value, OpMethodBody &body, 910 bool &shouldEmitSpace, bool &lastWasPunctuation) { 911 body << " p"; 912 913 // Don't insert a space for certain punctuation. 914 auto shouldPrintSpaceBeforeLiteral = [&] { 915 if (value.size() != 1 && value != "->") 916 return true; 917 if (lastWasPunctuation) 918 return !StringRef(">)}],").contains(value.front()); 919 return !StringRef("<>(){}[],").contains(value.front()); 920 }; 921 if (shouldEmitSpace && shouldPrintSpaceBeforeLiteral()) 922 body << " << \" \""; 923 body << " << \"" << value << "\";\n"; 924 925 // Insert a space after certain literals. 926 shouldEmitSpace = 927 value.size() != 1 || !StringRef("<({[").contains(value.front()); 928 lastWasPunctuation = !(value.front() == '_' || isalpha(value.front())); 929 } 930 931 /// Generate the C++ for an operand to a (*-)type directive. 932 static OpMethodBody &genTypeOperandPrinter(Element *arg, OpMethodBody &body) { 933 if (isa<OperandsDirective>(arg)) 934 return body << "getOperation()->getOperandTypes()"; 935 if (isa<ResultsDirective>(arg)) 936 return body << "getOperation()->getResultTypes()"; 937 auto *operand = dyn_cast<OperandVariable>(arg); 938 auto *var = operand ? operand->getVar() : cast<ResultVariable>(arg)->getVar(); 939 if (var->isVariadic()) 940 return body << var->name << "().getTypes()"; 941 if (var->isOptional()) 942 return body << llvm::formatv( 943 "({0}() ? ArrayRef<Type>({0}().getType()) : ArrayRef<Type>())", 944 var->name); 945 return body << "ArrayRef<Type>(" << var->name << "().getType())"; 946 } 947 948 /// Generate the code for printing the given element. 949 static void genElementPrinter(Element *element, OpMethodBody &body, 950 OperationFormat &fmt, Operator &op, 951 bool &shouldEmitSpace, bool &lastWasPunctuation) { 952 if (LiteralElement *literal = dyn_cast<LiteralElement>(element)) 953 return genLiteralPrinter(literal->getLiteral(), body, shouldEmitSpace, 954 lastWasPunctuation); 955 956 // Emit an optional group. 957 if (OptionalElement *optional = dyn_cast<OptionalElement>(element)) { 958 // Emit the check for the presence of the anchor element. 959 Element *anchor = optional->getAnchor(); 960 if (auto *operand = dyn_cast<OperandVariable>(anchor)) { 961 const NamedTypeConstraint *var = operand->getVar(); 962 if (var->isOptional()) 963 body << " if (" << var->name << "()) {\n"; 964 else if (var->isVariadic()) 965 body << " if (!" << var->name << "().empty()) {\n"; 966 } else { 967 body << " if (getAttr(\"" 968 << cast<AttributeVariable>(anchor)->getVar()->name << "\")) {\n"; 969 } 970 971 // Emit each of the elements. 972 for (Element &childElement : optional->getElements()) 973 genElementPrinter(&childElement, body, fmt, op, shouldEmitSpace, 974 lastWasPunctuation); 975 body << " }\n"; 976 return; 977 } 978 979 // Emit the attribute dictionary. 980 if (auto *attrDict = dyn_cast<AttrDictDirective>(element)) { 981 genAttrDictPrinter(fmt, op, body, attrDict->isWithKeyword()); 982 lastWasPunctuation = false; 983 return; 984 } 985 986 // Optionally insert a space before the next element. The AttrDict printer 987 // already adds a space as necessary. 988 if (shouldEmitSpace || !lastWasPunctuation) 989 body << " p << \" \";\n"; 990 lastWasPunctuation = false; 991 shouldEmitSpace = true; 992 993 if (auto *attr = dyn_cast<AttributeVariable>(element)) { 994 const NamedAttribute *var = attr->getVar(); 995 996 // If we are formatting as an enum, symbolize the attribute as a string. 997 if (canFormatEnumAttr(var)) { 998 const EnumAttr &enumAttr = cast<EnumAttr>(var->attr); 999 body << " p << \"\\\"\" << " << enumAttr.getSymbolToStringFnName() << "(" 1000 << var->name << "()) << \"\\\"\";\n"; 1001 return; 1002 } 1003 1004 // Elide the attribute type if it is buildable. 1005 if (attr->getTypeBuilder()) 1006 body << " p.printAttributeWithoutType(" << var->name << "Attr());\n"; 1007 else 1008 body << " p.printAttribute(" << var->name << "Attr());\n"; 1009 } else if (auto *operand = dyn_cast<OperandVariable>(element)) { 1010 if (operand->getVar()->isOptional()) { 1011 body << " if (Value value = " << operand->getVar()->name << "())\n" 1012 << " p << value;\n"; 1013 } else { 1014 body << " p << " << operand->getVar()->name << "();\n"; 1015 } 1016 } else if (auto *successor = dyn_cast<SuccessorVariable>(element)) { 1017 const NamedSuccessor *var = successor->getVar(); 1018 if (var->isVariadic()) 1019 body << " llvm::interleaveComma(" << var->name << "(), p);\n"; 1020 else 1021 body << " p << " << var->name << "();\n"; 1022 } else if (isa<OperandsDirective>(element)) { 1023 body << " p << getOperation()->getOperands();\n"; 1024 } else if (isa<SuccessorsDirective>(element)) { 1025 body << " llvm::interleaveComma(getOperation()->getSuccessors(), p);\n"; 1026 } else if (auto *dir = dyn_cast<TypeDirective>(element)) { 1027 body << " p << "; 1028 genTypeOperandPrinter(dir->getOperand(), body) << ";\n"; 1029 } else if (auto *dir = dyn_cast<FunctionalTypeDirective>(element)) { 1030 body << " p.printFunctionalType("; 1031 genTypeOperandPrinter(dir->getInputs(), body) << ", "; 1032 genTypeOperandPrinter(dir->getResults(), body) << ");\n"; 1033 } else { 1034 llvm_unreachable("unknown format element"); 1035 } 1036 } 1037 1038 void OperationFormat::genPrinter(Operator &op, OpClass &opClass) { 1039 auto &method = opClass.newMethod("void", "print", "OpAsmPrinter &p"); 1040 auto &body = method.body(); 1041 1042 // Emit the operation name, trimming the prefix if this is the standard 1043 // dialect. 1044 body << " p << \""; 1045 std::string opName = op.getOperationName(); 1046 if (op.getDialectName() == "std") 1047 body << StringRef(opName).drop_front(4); 1048 else 1049 body << opName; 1050 body << "\";\n"; 1051 1052 // Flags for if we should emit a space, and if the last element was 1053 // punctuation. 1054 bool shouldEmitSpace = true, lastWasPunctuation = false; 1055 for (auto &element : elements) 1056 genElementPrinter(element.get(), body, *this, op, shouldEmitSpace, 1057 lastWasPunctuation); 1058 } 1059 1060 //===----------------------------------------------------------------------===// 1061 // FormatLexer 1062 //===----------------------------------------------------------------------===// 1063 1064 namespace { 1065 /// This class represents a specific token in the input format. 1066 class Token { 1067 public: 1068 enum Kind { 1069 // Markers. 1070 eof, 1071 error, 1072 1073 // Tokens with no info. 1074 l_paren, 1075 r_paren, 1076 caret, 1077 comma, 1078 equal, 1079 question, 1080 1081 // Keywords. 1082 keyword_start, 1083 kw_attr_dict, 1084 kw_attr_dict_w_keyword, 1085 kw_functional_type, 1086 kw_operands, 1087 kw_results, 1088 kw_successors, 1089 kw_type, 1090 keyword_end, 1091 1092 // String valued tokens. 1093 identifier, 1094 literal, 1095 variable, 1096 }; 1097 Token(Kind kind, StringRef spelling) : kind(kind), spelling(spelling) {} 1098 1099 /// Return the bytes that make up this token. 1100 StringRef getSpelling() const { return spelling; } 1101 1102 /// Return the kind of this token. 1103 Kind getKind() const { return kind; } 1104 1105 /// Return a location for this token. 1106 llvm::SMLoc getLoc() const { 1107 return llvm::SMLoc::getFromPointer(spelling.data()); 1108 } 1109 1110 /// Return if this token is a keyword. 1111 bool isKeyword() const { return kind > keyword_start && kind < keyword_end; } 1112 1113 private: 1114 /// Discriminator that indicates the kind of token this is. 1115 Kind kind; 1116 1117 /// A reference to the entire token contents; this is always a pointer into 1118 /// a memory buffer owned by the source manager. 1119 StringRef spelling; 1120 }; 1121 1122 /// This class implements a simple lexer for operation assembly format strings. 1123 class FormatLexer { 1124 public: 1125 FormatLexer(llvm::SourceMgr &mgr, Operator &op); 1126 1127 /// Lex the next token and return it. 1128 Token lexToken(); 1129 1130 /// Emit an error to the lexer with the given location and message. 1131 Token emitError(llvm::SMLoc loc, const Twine &msg); 1132 Token emitError(const char *loc, const Twine &msg); 1133 1134 Token emitErrorAndNote(llvm::SMLoc loc, const Twine &msg, const Twine ¬e); 1135 1136 private: 1137 Token formToken(Token::Kind kind, const char *tokStart) { 1138 return Token(kind, StringRef(tokStart, curPtr - tokStart)); 1139 } 1140 1141 /// Return the next character in the stream. 1142 int getNextChar(); 1143 1144 /// Lex an identifier, literal, or variable. 1145 Token lexIdentifier(const char *tokStart); 1146 Token lexLiteral(const char *tokStart); 1147 Token lexVariable(const char *tokStart); 1148 1149 llvm::SourceMgr &srcMgr; 1150 Operator &op; 1151 StringRef curBuffer; 1152 const char *curPtr; 1153 }; 1154 } // end anonymous namespace 1155 1156 FormatLexer::FormatLexer(llvm::SourceMgr &mgr, Operator &op) 1157 : srcMgr(mgr), op(op) { 1158 curBuffer = srcMgr.getMemoryBuffer(mgr.getMainFileID())->getBuffer(); 1159 curPtr = curBuffer.begin(); 1160 } 1161 1162 Token FormatLexer::emitError(llvm::SMLoc loc, const Twine &msg) { 1163 srcMgr.PrintMessage(loc, llvm::SourceMgr::DK_Error, msg); 1164 llvm::SrcMgr.PrintMessage(op.getLoc()[0], llvm::SourceMgr::DK_Note, 1165 "in custom assembly format for this operation"); 1166 return formToken(Token::error, loc.getPointer()); 1167 } 1168 Token FormatLexer::emitErrorAndNote(llvm::SMLoc loc, const Twine &msg, 1169 const Twine ¬e) { 1170 srcMgr.PrintMessage(loc, llvm::SourceMgr::DK_Error, msg); 1171 llvm::SrcMgr.PrintMessage(op.getLoc()[0], llvm::SourceMgr::DK_Note, 1172 "in custom assembly format for this operation"); 1173 srcMgr.PrintMessage(loc, llvm::SourceMgr::DK_Note, note); 1174 return formToken(Token::error, loc.getPointer()); 1175 } 1176 Token FormatLexer::emitError(const char *loc, const Twine &msg) { 1177 return emitError(llvm::SMLoc::getFromPointer(loc), msg); 1178 } 1179 1180 int FormatLexer::getNextChar() { 1181 char curChar = *curPtr++; 1182 switch (curChar) { 1183 default: 1184 return (unsigned char)curChar; 1185 case 0: { 1186 // A nul character in the stream is either the end of the current buffer or 1187 // a random nul in the file. Disambiguate that here. 1188 if (curPtr - 1 != curBuffer.end()) 1189 return 0; 1190 1191 // Otherwise, return end of file. 1192 --curPtr; 1193 return EOF; 1194 } 1195 case '\n': 1196 case '\r': 1197 // Handle the newline character by ignoring it and incrementing the line 1198 // count. However, be careful about 'dos style' files with \n\r in them. 1199 // Only treat a \n\r or \r\n as a single line. 1200 if ((*curPtr == '\n' || (*curPtr == '\r')) && *curPtr != curChar) 1201 ++curPtr; 1202 return '\n'; 1203 } 1204 } 1205 1206 Token FormatLexer::lexToken() { 1207 const char *tokStart = curPtr; 1208 1209 // This always consumes at least one character. 1210 int curChar = getNextChar(); 1211 switch (curChar) { 1212 default: 1213 // Handle identifiers: [a-zA-Z_] 1214 if (isalpha(curChar) || curChar == '_') 1215 return lexIdentifier(tokStart); 1216 1217 // Unknown character, emit an error. 1218 return emitError(tokStart, "unexpected character"); 1219 case EOF: 1220 // Return EOF denoting the end of lexing. 1221 return formToken(Token::eof, tokStart); 1222 1223 // Lex punctuation. 1224 case '^': 1225 return formToken(Token::caret, tokStart); 1226 case ',': 1227 return formToken(Token::comma, tokStart); 1228 case '=': 1229 return formToken(Token::equal, tokStart); 1230 case '?': 1231 return formToken(Token::question, tokStart); 1232 case '(': 1233 return formToken(Token::l_paren, tokStart); 1234 case ')': 1235 return formToken(Token::r_paren, tokStart); 1236 1237 // Ignore whitespace characters. 1238 case 0: 1239 case ' ': 1240 case '\t': 1241 case '\n': 1242 return lexToken(); 1243 1244 case '`': 1245 return lexLiteral(tokStart); 1246 case '$': 1247 return lexVariable(tokStart); 1248 } 1249 } 1250 1251 Token FormatLexer::lexLiteral(const char *tokStart) { 1252 assert(curPtr[-1] == '`'); 1253 1254 // Lex a literal surrounded by ``. 1255 while (const char curChar = *curPtr++) { 1256 if (curChar == '`') 1257 return formToken(Token::literal, tokStart); 1258 } 1259 return emitError(curPtr - 1, "unexpected end of file in literal"); 1260 } 1261 1262 Token FormatLexer::lexVariable(const char *tokStart) { 1263 if (!isalpha(curPtr[0]) && curPtr[0] != '_') 1264 return emitError(curPtr - 1, "expected variable name"); 1265 1266 // Otherwise, consume the rest of the characters. 1267 while (isalnum(*curPtr) || *curPtr == '_') 1268 ++curPtr; 1269 return formToken(Token::variable, tokStart); 1270 } 1271 1272 Token FormatLexer::lexIdentifier(const char *tokStart) { 1273 // Match the rest of the identifier regex: [0-9a-zA-Z_\-]* 1274 while (isalnum(*curPtr) || *curPtr == '_' || *curPtr == '-') 1275 ++curPtr; 1276 1277 // Check to see if this identifier is a keyword. 1278 StringRef str(tokStart, curPtr - tokStart); 1279 Token::Kind kind = 1280 llvm::StringSwitch<Token::Kind>(str) 1281 .Case("attr-dict", Token::kw_attr_dict) 1282 .Case("attr-dict-with-keyword", Token::kw_attr_dict_w_keyword) 1283 .Case("functional-type", Token::kw_functional_type) 1284 .Case("operands", Token::kw_operands) 1285 .Case("results", Token::kw_results) 1286 .Case("successors", Token::kw_successors) 1287 .Case("type", Token::kw_type) 1288 .Default(Token::identifier); 1289 return Token(kind, str); 1290 } 1291 1292 //===----------------------------------------------------------------------===// 1293 // FormatParser 1294 //===----------------------------------------------------------------------===// 1295 1296 /// Function to find an element within the given range that has the same name as 1297 /// 'name'. 1298 template <typename RangeT> static auto findArg(RangeT &&range, StringRef name) { 1299 auto it = llvm::find_if(range, [=](auto &arg) { return arg.name == name; }); 1300 return it != range.end() ? &*it : nullptr; 1301 } 1302 1303 namespace { 1304 /// This class implements a parser for an instance of an operation assembly 1305 /// format. 1306 class FormatParser { 1307 public: 1308 FormatParser(llvm::SourceMgr &mgr, OperationFormat &format, Operator &op) 1309 : lexer(mgr, op), curToken(lexer.lexToken()), fmt(format), op(op), 1310 seenOperandTypes(op.getNumOperands()), 1311 seenResultTypes(op.getNumResults()) {} 1312 1313 /// Parse the operation assembly format. 1314 LogicalResult parse(); 1315 1316 private: 1317 /// This struct represents a type resolution instance. It includes a specific 1318 /// type as well as an optional transformer to apply to that type in order to 1319 /// properly resolve the type of a variable. 1320 struct TypeResolutionInstance { 1321 const NamedTypeConstraint *type; 1322 Optional<StringRef> transformer; 1323 }; 1324 1325 /// An iterator over the elements of a format group. 1326 using ElementsIterT = llvm::pointee_iterator< 1327 std::vector<std::unique_ptr<Element>>::const_iterator>; 1328 1329 /// Verify the state of operation attributes within the format. 1330 LogicalResult verifyAttributes(llvm::SMLoc loc); 1331 /// Verify the attribute elements at the back of the given stack of iterators. 1332 LogicalResult verifyAttributes( 1333 llvm::SMLoc loc, 1334 SmallVectorImpl<std::pair<ElementsIterT, ElementsIterT>> &iteratorStack); 1335 1336 /// Verify the state of operation operands within the format. 1337 LogicalResult 1338 verifyOperands(llvm::SMLoc loc, 1339 llvm::StringMap<TypeResolutionInstance> &variableTyResolver); 1340 1341 /// Verify the state of operation results within the format. 1342 LogicalResult 1343 verifyResults(llvm::SMLoc loc, 1344 llvm::StringMap<TypeResolutionInstance> &variableTyResolver); 1345 1346 /// Verify the state of operation successors within the format. 1347 LogicalResult verifySuccessors(llvm::SMLoc loc); 1348 1349 /// Given the values of an `AllTypesMatch` trait, check for inferable type 1350 /// resolution. 1351 void handleAllTypesMatchConstraint( 1352 ArrayRef<StringRef> values, 1353 llvm::StringMap<TypeResolutionInstance> &variableTyResolver); 1354 /// Check for inferable type resolution given all operands, and or results, 1355 /// have the same type. If 'includeResults' is true, the results also have the 1356 /// same type as all of the operands. 1357 void handleSameTypesConstraint( 1358 llvm::StringMap<TypeResolutionInstance> &variableTyResolver, 1359 bool includeResults); 1360 1361 /// Returns an argument with the given name that has been seen within the 1362 /// format. 1363 const NamedTypeConstraint *findSeenArg(StringRef name); 1364 1365 /// Parse a specific element. 1366 LogicalResult parseElement(std::unique_ptr<Element> &element, 1367 bool isTopLevel); 1368 LogicalResult parseVariable(std::unique_ptr<Element> &element, 1369 bool isTopLevel); 1370 LogicalResult parseDirective(std::unique_ptr<Element> &element, 1371 bool isTopLevel); 1372 LogicalResult parseLiteral(std::unique_ptr<Element> &element); 1373 LogicalResult parseOptional(std::unique_ptr<Element> &element, 1374 bool isTopLevel); 1375 LogicalResult parseOptionalChildElement( 1376 std::vector<std::unique_ptr<Element>> &childElements, 1377 SmallPtrSetImpl<const NamedTypeConstraint *> &seenVariables, 1378 Optional<unsigned> &anchorIdx); 1379 1380 /// Parse the various different directives. 1381 LogicalResult parseAttrDictDirective(std::unique_ptr<Element> &element, 1382 llvm::SMLoc loc, bool isTopLevel, 1383 bool withKeyword); 1384 LogicalResult parseFunctionalTypeDirective(std::unique_ptr<Element> &element, 1385 Token tok, bool isTopLevel); 1386 LogicalResult parseOperandsDirective(std::unique_ptr<Element> &element, 1387 llvm::SMLoc loc, bool isTopLevel); 1388 LogicalResult parseResultsDirective(std::unique_ptr<Element> &element, 1389 llvm::SMLoc loc, bool isTopLevel); 1390 LogicalResult parseSuccessorsDirective(std::unique_ptr<Element> &element, 1391 llvm::SMLoc loc, bool isTopLevel); 1392 LogicalResult parseTypeDirective(std::unique_ptr<Element> &element, Token tok, 1393 bool isTopLevel); 1394 LogicalResult parseTypeDirectiveOperand(std::unique_ptr<Element> &element); 1395 1396 //===--------------------------------------------------------------------===// 1397 // Lexer Utilities 1398 //===--------------------------------------------------------------------===// 1399 1400 /// Advance the current lexer onto the next token. 1401 void consumeToken() { 1402 assert(curToken.getKind() != Token::eof && 1403 curToken.getKind() != Token::error && 1404 "shouldn't advance past EOF or errors"); 1405 curToken = lexer.lexToken(); 1406 } 1407 LogicalResult parseToken(Token::Kind kind, const Twine &msg) { 1408 if (curToken.getKind() != kind) 1409 return emitError(curToken.getLoc(), msg); 1410 consumeToken(); 1411 return success(); 1412 } 1413 LogicalResult emitError(llvm::SMLoc loc, const Twine &msg) { 1414 lexer.emitError(loc, msg); 1415 return failure(); 1416 } 1417 LogicalResult emitErrorAndNote(llvm::SMLoc loc, const Twine &msg, 1418 const Twine ¬e) { 1419 lexer.emitErrorAndNote(loc, msg, note); 1420 return failure(); 1421 } 1422 1423 //===--------------------------------------------------------------------===// 1424 // Fields 1425 //===--------------------------------------------------------------------===// 1426 1427 FormatLexer lexer; 1428 Token curToken; 1429 OperationFormat &fmt; 1430 Operator &op; 1431 1432 // The following are various bits of format state used for verification 1433 // during parsing. 1434 bool hasAllOperands = false, hasAttrDict = false; 1435 bool hasAllSuccessors = false; 1436 llvm::SmallBitVector seenOperandTypes, seenResultTypes; 1437 llvm::DenseSet<const NamedTypeConstraint *> seenOperands; 1438 llvm::DenseSet<const NamedAttribute *> seenAttrs; 1439 llvm::DenseSet<const NamedSuccessor *> seenSuccessors; 1440 llvm::DenseSet<const NamedTypeConstraint *> optionalVariables; 1441 }; 1442 } // end anonymous namespace 1443 1444 LogicalResult FormatParser::parse() { 1445 llvm::SMLoc loc = curToken.getLoc(); 1446 1447 // Parse each of the format elements into the main format. 1448 while (curToken.getKind() != Token::eof) { 1449 std::unique_ptr<Element> element; 1450 if (failed(parseElement(element, /*isTopLevel=*/true))) 1451 return failure(); 1452 fmt.elements.push_back(std::move(element)); 1453 } 1454 1455 // Check that the attribute dictionary is in the format. 1456 if (!hasAttrDict) 1457 return emitError(loc, "'attr-dict' directive not found in " 1458 "custom assembly format"); 1459 1460 // Check for any type traits that we can use for inferring types. 1461 llvm::StringMap<TypeResolutionInstance> variableTyResolver; 1462 for (const OpTrait &trait : op.getTraits()) { 1463 const llvm::Record &def = trait.getDef(); 1464 if (def.isSubClassOf("AllTypesMatch")) { 1465 handleAllTypesMatchConstraint(def.getValueAsListOfStrings("values"), 1466 variableTyResolver); 1467 } else if (def.getName() == "SameTypeOperands") { 1468 handleSameTypesConstraint(variableTyResolver, /*includeResults=*/false); 1469 } else if (def.getName() == "SameOperandsAndResultType") { 1470 handleSameTypesConstraint(variableTyResolver, /*includeResults=*/true); 1471 } else if (def.isSubClassOf("TypesMatchWith")) { 1472 if (const auto *lhsArg = findSeenArg(def.getValueAsString("lhs"))) 1473 variableTyResolver[def.getValueAsString("rhs")] = { 1474 lhsArg, def.getValueAsString("transformer")}; 1475 } 1476 } 1477 1478 // Verify the state of the various operation components. 1479 if (failed(verifyAttributes(loc)) || 1480 failed(verifyResults(loc, variableTyResolver)) || 1481 failed(verifyOperands(loc, variableTyResolver)) || 1482 failed(verifySuccessors(loc))) 1483 return failure(); 1484 1485 // Check to see if we are formatting all of the operands. 1486 fmt.allOperands = llvm::any_of(fmt.elements, [](auto &elt) { 1487 return isa<OperandsDirective>(elt.get()); 1488 }); 1489 return success(); 1490 } 1491 1492 LogicalResult FormatParser::verifyAttributes(llvm::SMLoc loc) { 1493 // Check that there are no `:` literals after an attribute without a constant 1494 // type. The attribute grammar contains an optional trailing colon type, which 1495 // can lead to unexpected and generally unintended behavior. Given that, it is 1496 // better to just error out here instead. 1497 using ElementsIterT = llvm::pointee_iterator< 1498 std::vector<std::unique_ptr<Element>>::const_iterator>; 1499 SmallVector<std::pair<ElementsIterT, ElementsIterT>, 1> iteratorStack; 1500 iteratorStack.emplace_back(fmt.elements.begin(), fmt.elements.end()); 1501 while (!iteratorStack.empty()) 1502 if (failed(verifyAttributes(loc, iteratorStack))) 1503 return failure(); 1504 return success(); 1505 } 1506 /// Verify the attribute elements at the back of the given stack of iterators. 1507 LogicalResult FormatParser::verifyAttributes( 1508 llvm::SMLoc loc, 1509 SmallVectorImpl<std::pair<ElementsIterT, ElementsIterT>> &iteratorStack) { 1510 auto &stackIt = iteratorStack.back(); 1511 ElementsIterT &it = stackIt.first, e = stackIt.second; 1512 while (it != e) { 1513 Element *element = &*(it++); 1514 1515 // Traverse into optional groups. 1516 if (auto *optional = dyn_cast<OptionalElement>(element)) { 1517 auto elements = optional->getElements(); 1518 iteratorStack.emplace_back(elements.begin(), elements.end()); 1519 return success(); 1520 } 1521 1522 // We are checking for an attribute element followed by a `:`, so there is 1523 // no need to check the end. 1524 if (it == e && iteratorStack.size() == 1) 1525 break; 1526 1527 // Check for an attribute with a constant type builder, followed by a `:`. 1528 auto *prevAttr = dyn_cast<AttributeVariable>(element); 1529 if (!prevAttr || prevAttr->getTypeBuilder()) 1530 continue; 1531 1532 // Check the next iterator within the stack for literal elements. 1533 for (auto &nextItPair : iteratorStack) { 1534 ElementsIterT nextIt = nextItPair.first, nextE = nextItPair.second; 1535 for (; nextIt != nextE; ++nextIt) { 1536 // Skip any trailing optional groups or attribute dictionaries. 1537 if (isa<AttrDictDirective>(*nextIt) || isa<OptionalElement>(*nextIt)) 1538 continue; 1539 1540 // We are only interested in `:` literals. 1541 auto *literal = dyn_cast<LiteralElement>(&*nextIt); 1542 if (!literal || literal->getLiteral() != ":") 1543 break; 1544 1545 // TODO: Use the location of the literal element itself. 1546 return emitError( 1547 loc, llvm::formatv("format ambiguity caused by `:` literal found " 1548 "after attribute `{0}` which does not have " 1549 "a buildable type", 1550 prevAttr->getVar()->name)); 1551 } 1552 } 1553 } 1554 iteratorStack.pop_back(); 1555 return success(); 1556 } 1557 1558 LogicalResult FormatParser::verifyOperands( 1559 llvm::SMLoc loc, 1560 llvm::StringMap<TypeResolutionInstance> &variableTyResolver) { 1561 // Check that all of the operands are within the format, and their types can 1562 // be inferred. 1563 auto &buildableTypes = fmt.buildableTypes; 1564 for (unsigned i = 0, e = op.getNumOperands(); i != e; ++i) { 1565 NamedTypeConstraint &operand = op.getOperand(i); 1566 1567 // Check that the operand itself is in the format. 1568 if (!hasAllOperands && !seenOperands.count(&operand)) { 1569 return emitErrorAndNote(loc, 1570 "operand #" + Twine(i) + ", named '" + 1571 operand.name + "', not found", 1572 "suggest adding a '$" + operand.name + 1573 "' directive to the custom assembly format"); 1574 } 1575 1576 // Check that the operand type is in the format, or that it can be inferred. 1577 if (fmt.allOperandTypes || seenOperandTypes.test(i)) 1578 continue; 1579 1580 // Check to see if we can infer this type from another variable. 1581 auto varResolverIt = variableTyResolver.find(op.getOperand(i).name); 1582 if (varResolverIt != variableTyResolver.end()) { 1583 fmt.operandTypes[i].setVariable(varResolverIt->second.type, 1584 varResolverIt->second.transformer); 1585 continue; 1586 } 1587 1588 // Similarly to results, allow a custom builder for resolving the type if 1589 // we aren't using the 'operands' directive. 1590 Optional<StringRef> builder = operand.constraint.getBuilderCall(); 1591 if (!builder || (hasAllOperands && operand.isVariableLength())) { 1592 return emitErrorAndNote( 1593 loc, 1594 "type of operand #" + Twine(i) + ", named '" + operand.name + 1595 "', is not buildable and a buildable type cannot be inferred", 1596 "suggest adding a type constraint to the operation or adding a " 1597 "'type($" + 1598 operand.name + ")' directive to the " + "custom assembly format"); 1599 } 1600 auto it = buildableTypes.insert({*builder, buildableTypes.size()}); 1601 fmt.operandTypes[i].setBuilderIdx(it.first->second); 1602 } 1603 return success(); 1604 } 1605 1606 LogicalResult FormatParser::verifyResults( 1607 llvm::SMLoc loc, 1608 llvm::StringMap<TypeResolutionInstance> &variableTyResolver) { 1609 // If we format all of the types together, there is nothing to check. 1610 if (fmt.allResultTypes) 1611 return success(); 1612 1613 // Check that all of the result types can be inferred. 1614 auto &buildableTypes = fmt.buildableTypes; 1615 for (unsigned i = 0, e = op.getNumResults(); i != e; ++i) { 1616 if (seenResultTypes.test(i)) 1617 continue; 1618 1619 // Check to see if we can infer this type from another variable. 1620 auto varResolverIt = variableTyResolver.find(op.getResultName(i)); 1621 if (varResolverIt != variableTyResolver.end()) { 1622 fmt.resultTypes[i].setVariable(varResolverIt->second.type, 1623 varResolverIt->second.transformer); 1624 continue; 1625 } 1626 1627 // If the result is not variable length, allow for the case where the type 1628 // has a builder that we can use. 1629 NamedTypeConstraint &result = op.getResult(i); 1630 Optional<StringRef> builder = result.constraint.getBuilderCall(); 1631 if (!builder || result.isVariableLength()) { 1632 return emitErrorAndNote( 1633 loc, 1634 "type of result #" + Twine(i) + ", named '" + result.name + 1635 "', is not buildable and a buildable type cannot be inferred", 1636 "suggest adding a type constraint to the operation or adding a " 1637 "'type($" + 1638 result.name + ")' directive to the " + "custom assembly format"); 1639 } 1640 // Note in the format that this result uses the custom builder. 1641 auto it = buildableTypes.insert({*builder, buildableTypes.size()}); 1642 fmt.resultTypes[i].setBuilderIdx(it.first->second); 1643 } 1644 return success(); 1645 } 1646 1647 LogicalResult FormatParser::verifySuccessors(llvm::SMLoc loc) { 1648 // Check that all of the successors are within the format. 1649 if (hasAllSuccessors) 1650 return success(); 1651 1652 for (unsigned i = 0, e = op.getNumSuccessors(); i != e; ++i) { 1653 const NamedSuccessor &successor = op.getSuccessor(i); 1654 if (!seenSuccessors.count(&successor)) { 1655 return emitErrorAndNote(loc, 1656 "successor #" + Twine(i) + ", named '" + 1657 successor.name + "', not found", 1658 "suggest adding a '$" + successor.name + 1659 "' directive to the custom assembly format"); 1660 } 1661 } 1662 return success(); 1663 } 1664 1665 void FormatParser::handleAllTypesMatchConstraint( 1666 ArrayRef<StringRef> values, 1667 llvm::StringMap<TypeResolutionInstance> &variableTyResolver) { 1668 for (unsigned i = 0, e = values.size(); i != e; ++i) { 1669 // Check to see if this value matches a resolved operand or result type. 1670 const NamedTypeConstraint *arg = findSeenArg(values[i]); 1671 if (!arg) 1672 continue; 1673 1674 // Mark this value as the type resolver for the other variables. 1675 for (unsigned j = 0; j != i; ++j) 1676 variableTyResolver[values[j]] = {arg, llvm::None}; 1677 for (unsigned j = i + 1; j != e; ++j) 1678 variableTyResolver[values[j]] = {arg, llvm::None}; 1679 } 1680 } 1681 1682 void FormatParser::handleSameTypesConstraint( 1683 llvm::StringMap<TypeResolutionInstance> &variableTyResolver, 1684 bool includeResults) { 1685 const NamedTypeConstraint *resolver = nullptr; 1686 int resolvedIt = -1; 1687 1688 // Check to see if there is an operand or result to use for the resolution. 1689 if ((resolvedIt = seenOperandTypes.find_first()) != -1) 1690 resolver = &op.getOperand(resolvedIt); 1691 else if (includeResults && (resolvedIt = seenResultTypes.find_first()) != -1) 1692 resolver = &op.getResult(resolvedIt); 1693 else 1694 return; 1695 1696 // Set the resolvers for each operand and result. 1697 for (unsigned i = 0, e = op.getNumOperands(); i != e; ++i) 1698 if (!seenOperandTypes.test(i) && !op.getOperand(i).name.empty()) 1699 variableTyResolver[op.getOperand(i).name] = {resolver, llvm::None}; 1700 if (includeResults) { 1701 for (unsigned i = 0, e = op.getNumResults(); i != e; ++i) 1702 if (!seenResultTypes.test(i) && !op.getResultName(i).empty()) 1703 variableTyResolver[op.getResultName(i)] = {resolver, llvm::None}; 1704 } 1705 } 1706 1707 const NamedTypeConstraint *FormatParser::findSeenArg(StringRef name) { 1708 if (auto *arg = findArg(op.getOperands(), name)) 1709 return seenOperandTypes.test(arg - op.operand_begin()) ? arg : nullptr; 1710 if (auto *arg = findArg(op.getResults(), name)) 1711 return seenResultTypes.test(arg - op.result_begin()) ? arg : nullptr; 1712 return nullptr; 1713 } 1714 1715 LogicalResult FormatParser::parseElement(std::unique_ptr<Element> &element, 1716 bool isTopLevel) { 1717 // Directives. 1718 if (curToken.isKeyword()) 1719 return parseDirective(element, isTopLevel); 1720 // Literals. 1721 if (curToken.getKind() == Token::literal) 1722 return parseLiteral(element); 1723 // Optionals. 1724 if (curToken.getKind() == Token::l_paren) 1725 return parseOptional(element, isTopLevel); 1726 // Variables. 1727 if (curToken.getKind() == Token::variable) 1728 return parseVariable(element, isTopLevel); 1729 return emitError(curToken.getLoc(), 1730 "expected directive, literal, variable, or optional group"); 1731 } 1732 1733 LogicalResult FormatParser::parseVariable(std::unique_ptr<Element> &element, 1734 bool isTopLevel) { 1735 Token varTok = curToken; 1736 consumeToken(); 1737 1738 StringRef name = varTok.getSpelling().drop_front(); 1739 llvm::SMLoc loc = varTok.getLoc(); 1740 1741 // Check that the parsed argument is something actually registered on the 1742 // op. 1743 /// Attributes 1744 if (const NamedAttribute *attr = findArg(op.getAttributes(), name)) { 1745 if (isTopLevel && !seenAttrs.insert(attr).second) 1746 return emitError(loc, "attribute '" + name + "' is already bound"); 1747 element = std::make_unique<AttributeVariable>(attr); 1748 return success(); 1749 } 1750 /// Operands 1751 if (const NamedTypeConstraint *operand = findArg(op.getOperands(), name)) { 1752 if (isTopLevel) { 1753 if (hasAllOperands || !seenOperands.insert(operand).second) 1754 return emitError(loc, "operand '" + name + "' is already bound"); 1755 } 1756 element = std::make_unique<OperandVariable>(operand); 1757 return success(); 1758 } 1759 /// Results. 1760 if (const auto *result = findArg(op.getResults(), name)) { 1761 if (isTopLevel) 1762 return emitError(loc, "results can not be used at the top level"); 1763 element = std::make_unique<ResultVariable>(result); 1764 return success(); 1765 } 1766 /// Successors. 1767 if (const auto *successor = findArg(op.getSuccessors(), name)) { 1768 if (!isTopLevel) 1769 return emitError(loc, "successors can only be used at the top level"); 1770 if (hasAllSuccessors || !seenSuccessors.insert(successor).second) 1771 return emitError(loc, "successor '" + name + "' is already bound"); 1772 element = std::make_unique<SuccessorVariable>(successor); 1773 return success(); 1774 } 1775 return emitError( 1776 loc, "expected variable to refer to an argument, result, or successor"); 1777 } 1778 1779 LogicalResult FormatParser::parseDirective(std::unique_ptr<Element> &element, 1780 bool isTopLevel) { 1781 Token dirTok = curToken; 1782 consumeToken(); 1783 1784 switch (dirTok.getKind()) { 1785 case Token::kw_attr_dict: 1786 return parseAttrDictDirective(element, dirTok.getLoc(), isTopLevel, 1787 /*withKeyword=*/false); 1788 case Token::kw_attr_dict_w_keyword: 1789 return parseAttrDictDirective(element, dirTok.getLoc(), isTopLevel, 1790 /*withKeyword=*/true); 1791 case Token::kw_functional_type: 1792 return parseFunctionalTypeDirective(element, dirTok, isTopLevel); 1793 case Token::kw_operands: 1794 return parseOperandsDirective(element, dirTok.getLoc(), isTopLevel); 1795 case Token::kw_results: 1796 return parseResultsDirective(element, dirTok.getLoc(), isTopLevel); 1797 case Token::kw_successors: 1798 return parseSuccessorsDirective(element, dirTok.getLoc(), isTopLevel); 1799 case Token::kw_type: 1800 return parseTypeDirective(element, dirTok, isTopLevel); 1801 1802 default: 1803 llvm_unreachable("unknown directive token"); 1804 } 1805 } 1806 1807 LogicalResult FormatParser::parseLiteral(std::unique_ptr<Element> &element) { 1808 Token literalTok = curToken; 1809 consumeToken(); 1810 1811 // Check that the parsed literal is valid. 1812 StringRef value = literalTok.getSpelling().drop_front().drop_back(); 1813 if (!LiteralElement::isValidLiteral(value)) 1814 return emitError(literalTok.getLoc(), "expected valid literal"); 1815 1816 element = std::make_unique<LiteralElement>(value); 1817 return success(); 1818 } 1819 1820 LogicalResult FormatParser::parseOptional(std::unique_ptr<Element> &element, 1821 bool isTopLevel) { 1822 llvm::SMLoc curLoc = curToken.getLoc(); 1823 if (!isTopLevel) 1824 return emitError(curLoc, "optional groups can only be used as top-level " 1825 "elements"); 1826 consumeToken(); 1827 1828 // Parse the child elements for this optional group. 1829 std::vector<std::unique_ptr<Element>> elements; 1830 SmallPtrSet<const NamedTypeConstraint *, 8> seenVariables; 1831 Optional<unsigned> anchorIdx; 1832 do { 1833 if (failed(parseOptionalChildElement(elements, seenVariables, anchorIdx))) 1834 return failure(); 1835 } while (curToken.getKind() != Token::r_paren); 1836 consumeToken(); 1837 if (failed(parseToken(Token::question, "expected '?' after optional group"))) 1838 return failure(); 1839 1840 // The optional group is required to have an anchor. 1841 if (!anchorIdx) 1842 return emitError(curLoc, "optional group specified no anchor element"); 1843 1844 // The first element of the group must be one that can be parsed/printed in an 1845 // optional fashion. 1846 if (!isa<LiteralElement>(&*elements.front()) && 1847 !isa<OperandVariable>(&*elements.front())) 1848 return emitError(curLoc, "first element of an operand group must be a " 1849 "literal or operand"); 1850 1851 // After parsing all of the elements, ensure that all type directives refer 1852 // only to elements within the group. 1853 auto checkTypeOperand = [&](Element *typeEle) { 1854 auto *opVar = dyn_cast<OperandVariable>(typeEle); 1855 const NamedTypeConstraint *var = opVar ? opVar->getVar() : nullptr; 1856 if (!seenVariables.count(var)) 1857 return emitError(curLoc, "type directive can only refer to variables " 1858 "within the optional group"); 1859 return success(); 1860 }; 1861 for (auto &ele : elements) { 1862 if (auto *typeEle = dyn_cast<TypeDirective>(ele.get())) { 1863 if (failed(checkTypeOperand(typeEle->getOperand()))) 1864 return failure(); 1865 } else if (auto *typeEle = dyn_cast<FunctionalTypeDirective>(ele.get())) { 1866 if (failed(checkTypeOperand(typeEle->getInputs())) || 1867 failed(checkTypeOperand(typeEle->getResults()))) 1868 return failure(); 1869 } 1870 } 1871 1872 optionalVariables.insert(seenVariables.begin(), seenVariables.end()); 1873 element = std::make_unique<OptionalElement>(std::move(elements), *anchorIdx); 1874 return success(); 1875 } 1876 1877 LogicalResult FormatParser::parseOptionalChildElement( 1878 std::vector<std::unique_ptr<Element>> &childElements, 1879 SmallPtrSetImpl<const NamedTypeConstraint *> &seenVariables, 1880 Optional<unsigned> &anchorIdx) { 1881 llvm::SMLoc childLoc = curToken.getLoc(); 1882 childElements.push_back({}); 1883 if (failed(parseElement(childElements.back(), /*isTopLevel=*/true))) 1884 return failure(); 1885 1886 // Check to see if this element is the anchor of the optional group. 1887 bool isAnchor = curToken.getKind() == Token::caret; 1888 if (isAnchor) { 1889 if (anchorIdx) 1890 return emitError(childLoc, "only one element can be marked as the anchor " 1891 "of an optional group"); 1892 anchorIdx = childElements.size() - 1; 1893 consumeToken(); 1894 } 1895 1896 return TypeSwitch<Element *, LogicalResult>(childElements.back().get()) 1897 // All attributes can be within the optional group, but only optional 1898 // attributes can be the anchor. 1899 .Case([&](AttributeVariable *attrEle) { 1900 if (isAnchor && !attrEle->getVar()->attr.isOptional()) 1901 return emitError(childLoc, "only optional attributes can be used to " 1902 "anchor an optional group"); 1903 return success(); 1904 }) 1905 // Only optional-like(i.e. variadic) operands can be within an optional 1906 // group. 1907 .Case<OperandVariable>([&](OperandVariable *ele) { 1908 if (!ele->getVar()->isVariableLength()) 1909 return emitError(childLoc, "only variable length operands can be " 1910 "used within an optional group"); 1911 seenVariables.insert(ele->getVar()); 1912 return success(); 1913 }) 1914 // Literals and type directives may be used, but they can't anchor the 1915 // group. 1916 .Case<LiteralElement, TypeDirective, FunctionalTypeDirective>( 1917 [&](Element *) { 1918 if (isAnchor) 1919 return emitError(childLoc, "only variables can be used to anchor " 1920 "an optional group"); 1921 return success(); 1922 }) 1923 .Default([&](Element *) { 1924 return emitError(childLoc, "only literals, types, and variables can be " 1925 "used within an optional group"); 1926 }); 1927 } 1928 1929 LogicalResult 1930 FormatParser::parseAttrDictDirective(std::unique_ptr<Element> &element, 1931 llvm::SMLoc loc, bool isTopLevel, 1932 bool withKeyword) { 1933 if (!isTopLevel) 1934 return emitError(loc, "'attr-dict' directive can only be used as a " 1935 "top-level directive"); 1936 if (hasAttrDict) 1937 return emitError(loc, "'attr-dict' directive has already been seen"); 1938 1939 hasAttrDict = true; 1940 element = std::make_unique<AttrDictDirective>(withKeyword); 1941 return success(); 1942 } 1943 1944 LogicalResult 1945 FormatParser::parseFunctionalTypeDirective(std::unique_ptr<Element> &element, 1946 Token tok, bool isTopLevel) { 1947 llvm::SMLoc loc = tok.getLoc(); 1948 if (!isTopLevel) 1949 return emitError( 1950 loc, "'functional-type' is only valid as a top-level directive"); 1951 1952 // Parse the main operand. 1953 std::unique_ptr<Element> inputs, results; 1954 if (failed(parseToken(Token::l_paren, "expected '(' before argument list")) || 1955 failed(parseTypeDirectiveOperand(inputs)) || 1956 failed(parseToken(Token::comma, "expected ',' after inputs argument")) || 1957 failed(parseTypeDirectiveOperand(results)) || 1958 failed(parseToken(Token::r_paren, "expected ')' after argument list"))) 1959 return failure(); 1960 element = std::make_unique<FunctionalTypeDirective>(std::move(inputs), 1961 std::move(results)); 1962 return success(); 1963 } 1964 1965 LogicalResult 1966 FormatParser::parseOperandsDirective(std::unique_ptr<Element> &element, 1967 llvm::SMLoc loc, bool isTopLevel) { 1968 if (isTopLevel && (hasAllOperands || !seenOperands.empty())) 1969 return emitError(loc, "'operands' directive creates overlap in format"); 1970 hasAllOperands = true; 1971 element = std::make_unique<OperandsDirective>(); 1972 return success(); 1973 } 1974 1975 LogicalResult 1976 FormatParser::parseResultsDirective(std::unique_ptr<Element> &element, 1977 llvm::SMLoc loc, bool isTopLevel) { 1978 if (isTopLevel) 1979 return emitError(loc, "'results' directive can not be used as a " 1980 "top-level directive"); 1981 element = std::make_unique<ResultsDirective>(); 1982 return success(); 1983 } 1984 1985 LogicalResult 1986 FormatParser::parseSuccessorsDirective(std::unique_ptr<Element> &element, 1987 llvm::SMLoc loc, bool isTopLevel) { 1988 if (!isTopLevel) 1989 return emitError(loc, 1990 "'successors' is only valid as a top-level directive"); 1991 if (hasAllSuccessors || !seenSuccessors.empty()) 1992 return emitError(loc, "'successors' directive creates overlap in format"); 1993 hasAllSuccessors = true; 1994 element = std::make_unique<SuccessorsDirective>(); 1995 return success(); 1996 } 1997 1998 LogicalResult 1999 FormatParser::parseTypeDirective(std::unique_ptr<Element> &element, Token tok, 2000 bool isTopLevel) { 2001 llvm::SMLoc loc = tok.getLoc(); 2002 if (!isTopLevel) 2003 return emitError(loc, "'type' is only valid as a top-level directive"); 2004 2005 std::unique_ptr<Element> operand; 2006 if (failed(parseToken(Token::l_paren, "expected '(' before argument list")) || 2007 failed(parseTypeDirectiveOperand(operand)) || 2008 failed(parseToken(Token::r_paren, "expected ')' after argument list"))) 2009 return failure(); 2010 element = std::make_unique<TypeDirective>(std::move(operand)); 2011 return success(); 2012 } 2013 2014 LogicalResult 2015 FormatParser::parseTypeDirectiveOperand(std::unique_ptr<Element> &element) { 2016 llvm::SMLoc loc = curToken.getLoc(); 2017 if (failed(parseElement(element, /*isTopLevel=*/false))) 2018 return failure(); 2019 if (isa<LiteralElement>(element.get())) 2020 return emitError( 2021 loc, "'type' directive operand expects variable or directive operand"); 2022 2023 if (auto *var = dyn_cast<OperandVariable>(element.get())) { 2024 unsigned opIdx = var->getVar() - op.operand_begin(); 2025 if (fmt.allOperandTypes || seenOperandTypes.test(opIdx)) 2026 return emitError(loc, "'type' of '" + var->getVar()->name + 2027 "' is already bound"); 2028 seenOperandTypes.set(opIdx); 2029 } else if (auto *var = dyn_cast<ResultVariable>(element.get())) { 2030 unsigned resIdx = var->getVar() - op.result_begin(); 2031 if (fmt.allResultTypes || seenResultTypes.test(resIdx)) 2032 return emitError(loc, "'type' of '" + var->getVar()->name + 2033 "' is already bound"); 2034 seenResultTypes.set(resIdx); 2035 } else if (isa<OperandsDirective>(&*element)) { 2036 if (fmt.allOperandTypes || seenOperandTypes.any()) 2037 return emitError(loc, "'operands' 'type' is already bound"); 2038 fmt.allOperandTypes = true; 2039 } else if (isa<ResultsDirective>(&*element)) { 2040 if (fmt.allResultTypes || seenResultTypes.any()) 2041 return emitError(loc, "'results' 'type' is already bound"); 2042 fmt.allResultTypes = true; 2043 } else { 2044 return emitError(loc, "invalid argument to 'type' directive"); 2045 } 2046 return success(); 2047 } 2048 2049 //===----------------------------------------------------------------------===// 2050 // Interface 2051 //===----------------------------------------------------------------------===// 2052 2053 void mlir::tblgen::generateOpFormat(const Operator &constOp, OpClass &opClass) { 2054 // TODO(riverriddle) Operator doesn't expose all necessary functionality via 2055 // the const interface. 2056 Operator &op = const_cast<Operator &>(constOp); 2057 if (!op.hasAssemblyFormat()) 2058 return; 2059 2060 // Parse the format description. 2061 llvm::SourceMgr mgr; 2062 mgr.AddNewSourceBuffer( 2063 llvm::MemoryBuffer::getMemBuffer(op.getAssemblyFormat()), llvm::SMLoc()); 2064 OperationFormat format(op); 2065 if (failed(FormatParser(mgr, format, op).parse())) { 2066 // Exit the process if format errors are treated as fatal. 2067 if (formatErrorIsFatal) { 2068 // Invoke the interrupt handlers to run the file cleanup handlers. 2069 llvm::sys::RunInterruptHandlers(); 2070 std::exit(1); 2071 } 2072 return; 2073 } 2074 2075 // Generate the printer and parser based on the parsed format. 2076 format.genParser(op, opClass); 2077 format.genPrinter(op, opClass); 2078 } 2079