1 //===- OpInterfacesGen.cpp - MLIR op interface utility 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 // OpInterfacesGen generates definitions for operation interfaces.
10 //
11 //===----------------------------------------------------------------------===//
12 
13 #include "DocGenUtilities.h"
14 #include "mlir/TableGen/Format.h"
15 #include "mlir/TableGen/GenInfo.h"
16 #include "mlir/TableGen/OpInterfaces.h"
17 #include "llvm/ADT/SmallVector.h"
18 #include "llvm/ADT/StringExtras.h"
19 #include "llvm/Support/FormatVariadic.h"
20 #include "llvm/Support/raw_ostream.h"
21 #include "llvm/TableGen/Error.h"
22 #include "llvm/TableGen/Record.h"
23 #include "llvm/TableGen/TableGenBackend.h"
24 
25 using namespace llvm;
26 using namespace mlir;
27 using mlir::tblgen::OpInterface;
28 using mlir::tblgen::OpInterfaceMethod;
29 
30 // Emit the method name and argument list for the given method. If
31 // 'addOperationArg' is true, then an Operation* argument is added to the
32 // beginning of the argument list.
33 static void emitMethodNameAndArgs(const OpInterfaceMethod &method,
34                                   raw_ostream &os, bool addOperationArg) {
35   os << method.getName() << '(';
36   if (addOperationArg)
37     os << "Operation *tablegen_opaque_op" << (method.arg_empty() ? "" : ", ");
38   llvm::interleaveComma(method.getArguments(), os,
39                         [&](const OpInterfaceMethod::Argument &arg) {
40                           os << arg.type << " " << arg.name;
41                         });
42   os << ')';
43 }
44 
45 // Get an array of all OpInterface definitions but exclude those subclassing
46 // "DeclareOpInterfaceMethods".
47 static std::vector<Record *>
48 getAllOpInterfaceDefinitions(const RecordKeeper &recordKeeper) {
49   std::vector<Record *> defs =
50       recordKeeper.getAllDerivedDefinitions("OpInterface");
51 
52   llvm::erase_if(defs, [](const Record *def) {
53     return def->isSubClassOf("DeclareOpInterfaceMethods");
54   });
55   return defs;
56 }
57 
58 //===----------------------------------------------------------------------===//
59 // GEN: Interface definitions
60 //===----------------------------------------------------------------------===//
61 
62 static void emitInterfaceDef(OpInterface &interface, raw_ostream &os) {
63   StringRef interfaceName = interface.getName();
64 
65   // Insert the method definitions.
66   for (auto &method : interface.getMethods()) {
67     os << method.getReturnType() << " " << interfaceName << "::";
68     emitMethodNameAndArgs(method, os, /*addOperationArg=*/false);
69 
70     // Forward to the method on the concrete operation type.
71     os << " {\n      return getImpl()->" << method.getName() << '(';
72     if (!method.isStatic())
73       os << "getOperation()" << (method.arg_empty() ? "" : ", ");
74     llvm::interleaveComma(
75         method.getArguments(), os,
76         [&](const OpInterfaceMethod::Argument &arg) { os << arg.name; });
77     os << ");\n  }\n";
78   }
79 }
80 
81 static bool emitInterfaceDefs(const RecordKeeper &recordKeeper,
82                               raw_ostream &os) {
83   llvm::emitSourceFileHeader("Operation Interface Definitions", os);
84 
85   for (const auto *def : getAllOpInterfaceDefinitions(recordKeeper)) {
86     OpInterface interface(def);
87     emitInterfaceDef(interface, os);
88   }
89   return false;
90 }
91 
92 //===----------------------------------------------------------------------===//
93 // GEN: Interface declarations
94 //===----------------------------------------------------------------------===//
95 
96 static void emitConceptDecl(OpInterface &interface, raw_ostream &os) {
97   os << "  class Concept {\n"
98      << "  public:\n"
99      << "    virtual ~Concept() = default;\n";
100 
101   // Insert each of the pure virtual concept methods.
102   for (auto &method : interface.getMethods()) {
103     os << "    virtual " << method.getReturnType() << " ";
104     emitMethodNameAndArgs(method, os, /*addOperationArg=*/!method.isStatic());
105     os << " = 0;\n";
106   }
107   os << "  };\n";
108 }
109 
110 static void emitModelDecl(OpInterface &interface, raw_ostream &os) {
111   os << "  template<typename ConcreteOp>\n";
112   os << "  class Model : public Concept {\npublic:\n";
113 
114   // Insert each of the virtual method overrides.
115   for (auto &method : interface.getMethods()) {
116     os << "    " << method.getReturnType() << " ";
117     emitMethodNameAndArgs(method, os, /*addOperationArg=*/!method.isStatic());
118     os << " final {\n";
119 
120     // Provide a definition of the concrete op if this is non static.
121     if (!method.isStatic()) {
122       os << "      auto op = llvm::cast<ConcreteOp>(tablegen_opaque_op);\n"
123          << "      (void)op;\n";
124     }
125 
126     // Check for a provided body to the function.
127     if (auto body = method.getBody()) {
128       os << body << "\n    }\n";
129       continue;
130     }
131 
132     // Forward to the method on the concrete operation type.
133     os << "      return " << (method.isStatic() ? "ConcreteOp::" : "op.");
134 
135     // Add the arguments to the call.
136     os << method.getName() << '(';
137     llvm::interleaveComma(
138         method.getArguments(), os,
139         [&](const OpInterfaceMethod::Argument &arg) { os << arg.name; });
140     os << ");\n    }\n";
141   }
142   os << "  };\n";
143 }
144 
145 static void emitTraitDecl(OpInterface &interface, raw_ostream &os,
146                           StringRef interfaceName,
147                           StringRef interfaceTraitsName) {
148   os << "  template <typename ConcreteOp>\n  "
149      << llvm::formatv("struct Trait : public OpInterface<{0},"
150                       " detail::{1}>::Trait<ConcreteOp> {{\n",
151                       interfaceName, interfaceTraitsName);
152 
153   // Insert the default implementation for any methods.
154   for (auto &method : interface.getMethods()) {
155     // Flag interface methods named verifyTrait.
156     if (method.getName() == "verifyTrait")
157       PrintFatalError(
158           formatv("'verifyTrait' method cannot be specified as interface "
159                   "method for '{0}'; set 'verify' on OpInterfaceTrait instead",
160                   interfaceName));
161     auto defaultImpl = method.getDefaultImplementation();
162     if (!defaultImpl)
163       continue;
164 
165     os << "  " << (method.isStatic() ? "static " : "") << method.getReturnType()
166        << " ";
167     emitMethodNameAndArgs(method, os, /*addOperationArg=*/false);
168     os << " {\n" << defaultImpl.getValue() << "  }\n";
169   }
170 
171   tblgen::FmtContext traitCtx;
172   traitCtx.withOp("op");
173   if (auto verify = interface.getVerify()) {
174     os << "  static LogicalResult verifyTrait(Operation* op) {\n"
175        << std::string(tblgen::tgfmt(*verify, &traitCtx)) << "\n  }\n";
176   }
177 
178   os << "  };\n";
179 }
180 
181 static void emitInterfaceDecl(OpInterface &interface, raw_ostream &os) {
182   StringRef interfaceName = interface.getName();
183   auto interfaceTraitsName = (interfaceName + "InterfaceTraits").str();
184 
185   // Emit the traits struct containing the concept and model declarations.
186   os << "namespace detail {\n"
187      << "struct " << interfaceTraitsName << " {\n";
188   emitConceptDecl(interface, os);
189   emitModelDecl(interface, os);
190   os << "};\n} // end namespace detail\n";
191 
192   // Emit the main interface class declaration.
193   os << llvm::formatv("class {0} : public OpInterface<{1}, detail::{2}> {\n"
194                       "public:\n"
195                       "  using OpInterface<{1}, detail::{2}>::OpInterface;\n",
196                       interfaceName, interfaceName, interfaceTraitsName);
197 
198   // Emit the derived trait for the interface.
199   emitTraitDecl(interface, os, interfaceName, interfaceTraitsName);
200 
201   // Insert the method declarations.
202   for (auto &method : interface.getMethods()) {
203     os << "  " << method.getReturnType() << " ";
204     emitMethodNameAndArgs(method, os, /*addOperationArg=*/false);
205     os << ";\n";
206   }
207 
208   // Emit any extra declarations.
209   if (Optional<StringRef> extraDecls = interface.getExtraClassDeclaration())
210     os << *extraDecls << "\n";
211 
212   os << "};\n";
213 }
214 
215 static bool emitInterfaceDecls(const RecordKeeper &recordKeeper,
216                                raw_ostream &os) {
217   llvm::emitSourceFileHeader("Operation Interface Declarations", os);
218 
219   for (const auto *def : getAllOpInterfaceDefinitions(recordKeeper)) {
220     OpInterface interface(def);
221     emitInterfaceDecl(interface, os);
222   }
223   return false;
224 }
225 
226 //===----------------------------------------------------------------------===//
227 // GEN: Interface documentation
228 //===----------------------------------------------------------------------===//
229 
230 /// Emit a string corresponding to a C++ type, followed by a space if necessary.
231 static raw_ostream &emitCPPType(StringRef type, raw_ostream &os) {
232   type = type.trim();
233   os << type;
234   if (type.back() != '&' && type.back() != '*')
235     os << " ";
236   return os;
237 }
238 
239 static void emitInterfaceDoc(const Record &interfaceDef, raw_ostream &os) {
240   OpInterface interface(&interfaceDef);
241 
242   // Emit the interface name followed by the description.
243   os << "## " << interface.getName() << " (" << interfaceDef.getName() << ")";
244   if (auto description = interface.getDescription())
245     mlir::tblgen::emitDescription(*description, os);
246 
247   // Emit the methods required by the interface.
248   os << "\n### Methods:\n";
249   for (const auto &method : interface.getMethods()) {
250     // Emit the method name.
251     os << "#### `" << method.getName() << "`\n\n```c++\n";
252 
253     // Emit the method signature.
254     if (method.isStatic())
255       os << "static ";
256     emitCPPType(method.getReturnType(), os) << method.getName() << '(';
257     llvm::interleaveComma(method.getArguments(), os,
258                           [&](const OpInterfaceMethod::Argument &arg) {
259                             emitCPPType(arg.type, os) << arg.name;
260                           });
261     os << ");\n```\n";
262 
263     // Emit the description.
264     if (auto description = method.getDescription())
265       mlir::tblgen::emitDescription(*description, os);
266 
267     // If the body is not provided, this method must be provided by the
268     // operation.
269     if (!method.getBody())
270       os << "\nNOTE: This method *must* be implemented by the operation.\n\n";
271   }
272 }
273 
274 static bool emitInterfaceDocs(const RecordKeeper &recordKeeper,
275                               raw_ostream &os) {
276   os << "<!-- Autogenerated by mlir-tblgen; don't manually edit -->\n";
277   os << "# Operation Interface definition\n";
278 
279   for (const auto *def : getAllOpInterfaceDefinitions(recordKeeper))
280     emitInterfaceDoc(*def, os);
281   return false;
282 }
283 
284 //===----------------------------------------------------------------------===//
285 // GEN: Interface registration hooks
286 //===----------------------------------------------------------------------===//
287 
288 // Registers the operation interface generator to mlir-tblgen.
289 static mlir::GenRegistration
290     genInterfaceDecls("gen-op-interface-decls",
291                       "Generate op interface declarations",
292                       [](const RecordKeeper &records, raw_ostream &os) {
293                         return emitInterfaceDecls(records, os);
294                       });
295 
296 // Registers the operation interface generator to mlir-tblgen.
297 static mlir::GenRegistration
298     genInterfaceDefs("gen-op-interface-defs",
299                      "Generate op interface definitions",
300                      [](const RecordKeeper &records, raw_ostream &os) {
301                        return emitInterfaceDefs(records, os);
302                      });
303 
304 // Registers the operation interface document generator to mlir-tblgen.
305 static mlir::GenRegistration
306     genInterfaceDocs("gen-op-interface-doc",
307                      "Generate op interface documentation",
308                      [](const RecordKeeper &records, raw_ostream &os) {
309                        return emitInterfaceDocs(records, os);
310                      });
311