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 &note);
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 &note) {
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 &note) {
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