1 //===- Diagnostics.cpp - MLIR Diagnostics ---------------------------------===//
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 "mlir/IR/Diagnostics.h"
10 #include "mlir/IR/Attributes.h"
11 #include "mlir/IR/Identifier.h"
12 #include "mlir/IR/Location.h"
13 #include "mlir/IR/MLIRContext.h"
14 #include "mlir/IR/Operation.h"
15 #include "mlir/IR/Types.h"
16 #include "llvm/ADT/MapVector.h"
17 #include "llvm/ADT/SmallString.h"
18 #include "llvm/ADT/StringMap.h"
19 #include "llvm/Support/Mutex.h"
20 #include "llvm/Support/PrettyStackTrace.h"
21 #include "llvm/Support/Regex.h"
22 #include "llvm/Support/Signals.h"
23 #include "llvm/Support/SourceMgr.h"
24 #include "llvm/Support/raw_ostream.h"
25 
26 using namespace mlir;
27 using namespace mlir::detail;
28 
29 //===----------------------------------------------------------------------===//
30 // DiagnosticArgument
31 //===----------------------------------------------------------------------===//
32 
33 /// Construct from an Attribute.
34 DiagnosticArgument::DiagnosticArgument(Attribute attr)
35     : kind(DiagnosticArgumentKind::Attribute),
36       opaqueVal(reinterpret_cast<intptr_t>(attr.getAsOpaquePointer())) {}
37 
38 /// Construct from a Type.
39 DiagnosticArgument::DiagnosticArgument(Type val)
40     : kind(DiagnosticArgumentKind::Type),
41       opaqueVal(reinterpret_cast<intptr_t>(val.getAsOpaquePointer())) {}
42 
43 /// Returns this argument as an Attribute.
44 Attribute DiagnosticArgument::getAsAttribute() const {
45   assert(getKind() == DiagnosticArgumentKind::Attribute);
46   return Attribute::getFromOpaquePointer(
47       reinterpret_cast<const void *>(opaqueVal));
48 }
49 
50 /// Returns this argument as a Type.
51 Type DiagnosticArgument::getAsType() const {
52   assert(getKind() == DiagnosticArgumentKind::Type);
53   return Type::getFromOpaquePointer(reinterpret_cast<const void *>(opaqueVal));
54 }
55 
56 /// Outputs this argument to a stream.
57 void DiagnosticArgument::print(raw_ostream &os) const {
58   switch (kind) {
59   case DiagnosticArgumentKind::Attribute:
60     os << getAsAttribute();
61     break;
62   case DiagnosticArgumentKind::Double:
63     os << getAsDouble();
64     break;
65   case DiagnosticArgumentKind::Integer:
66     os << getAsInteger();
67     break;
68   case DiagnosticArgumentKind::String:
69     os << getAsString();
70     break;
71   case DiagnosticArgumentKind::Type:
72     os << '\'' << getAsType() << '\'';
73     break;
74   case DiagnosticArgumentKind::Unsigned:
75     os << getAsUnsigned();
76     break;
77   }
78 }
79 
80 //===----------------------------------------------------------------------===//
81 // Diagnostic
82 //===----------------------------------------------------------------------===//
83 
84 /// Convert a Twine to a StringRef. Memory used for generating the StringRef is
85 /// stored in 'strings'.
86 static StringRef twineToStrRef(const Twine &val,
87                                std::vector<std::unique_ptr<char[]>> &strings) {
88   // Allocate memory to hold this string.
89   SmallString<64> data;
90   auto strRef = val.toStringRef(data);
91   strings.push_back(std::unique_ptr<char[]>(new char[strRef.size()]));
92   memcpy(&strings.back()[0], strRef.data(), strRef.size());
93 
94   // Return a reference to the new string.
95   return StringRef(&strings.back()[0], strRef.size());
96 }
97 
98 /// Stream in a Twine argument.
99 Diagnostic &Diagnostic::operator<<(char val) { return *this << Twine(val); }
100 Diagnostic &Diagnostic::operator<<(const Twine &val) {
101   arguments.push_back(DiagnosticArgument(twineToStrRef(val, strings)));
102   return *this;
103 }
104 Diagnostic &Diagnostic::operator<<(Twine &&val) {
105   arguments.push_back(DiagnosticArgument(twineToStrRef(val, strings)));
106   return *this;
107 }
108 
109 /// Stream in an Identifier.
110 Diagnostic &Diagnostic::operator<<(Identifier val) {
111   // An identifier is stored in the context, so we don't need to worry about the
112   // lifetime of its data.
113   arguments.push_back(DiagnosticArgument(val.strref()));
114   return *this;
115 }
116 
117 /// Stream in an OperationName.
118 Diagnostic &Diagnostic::operator<<(OperationName val) {
119   // An OperationName is stored in the context, so we don't need to worry about
120   // the lifetime of its data.
121   arguments.push_back(DiagnosticArgument(val.getStringRef()));
122   return *this;
123 }
124 
125 /// Stream in an Operation.
126 Diagnostic &Diagnostic::operator<<(Operation &val) {
127   std::string str;
128   llvm::raw_string_ostream os(str);
129   os << val;
130   return *this << os.str();
131 }
132 
133 /// Outputs this diagnostic to a stream.
134 void Diagnostic::print(raw_ostream &os) const {
135   for (auto &arg : getArguments())
136     arg.print(os);
137 }
138 
139 /// Convert the diagnostic to a string.
140 std::string Diagnostic::str() const {
141   std::string str;
142   llvm::raw_string_ostream os(str);
143   print(os);
144   return os.str();
145 }
146 
147 /// Attaches a note to this diagnostic. A new location may be optionally
148 /// provided, if not, then the location defaults to the one specified for this
149 /// diagnostic. Notes may not be attached to other notes.
150 Diagnostic &Diagnostic::attachNote(Optional<Location> noteLoc) {
151   // We don't allow attaching notes to notes.
152   assert(severity != DiagnosticSeverity::Note &&
153          "cannot attach a note to a note");
154 
155   // If a location wasn't provided then reuse our location.
156   if (!noteLoc)
157     noteLoc = loc;
158 
159   /// Append and return a new note.
160   notes.push_back(
161       std::make_unique<Diagnostic>(*noteLoc, DiagnosticSeverity::Note));
162   return *notes.back();
163 }
164 
165 /// Allow a diagnostic to be converted to 'failure'.
166 Diagnostic::operator LogicalResult() const { return failure(); }
167 
168 //===----------------------------------------------------------------------===//
169 // InFlightDiagnostic
170 //===----------------------------------------------------------------------===//
171 
172 /// Allow an inflight diagnostic to be converted to 'failure', otherwise
173 /// 'success' if this is an empty diagnostic.
174 InFlightDiagnostic::operator LogicalResult() const {
175   return failure(isActive());
176 }
177 
178 /// Reports the diagnostic to the engine.
179 void InFlightDiagnostic::report() {
180   // If this diagnostic is still inflight and it hasn't been abandoned, then
181   // report it.
182   if (isInFlight()) {
183     owner->emit(std::move(*impl));
184     owner = nullptr;
185   }
186   impl.reset();
187 }
188 
189 /// Abandons this diagnostic.
190 void InFlightDiagnostic::abandon() { owner = nullptr; }
191 
192 //===----------------------------------------------------------------------===//
193 // DiagnosticEngineImpl
194 //===----------------------------------------------------------------------===//
195 
196 namespace mlir {
197 namespace detail {
198 struct DiagnosticEngineImpl {
199   /// Emit a diagnostic using the registered issue handle if present, or with
200   /// the default behavior if not.
201   void emit(Diagnostic diag);
202 
203   /// A mutex to ensure that diagnostics emission is thread-safe.
204   llvm::sys::SmartMutex<true> mutex;
205 
206   /// These are the handlers used to report diagnostics.
207   llvm::SmallMapVector<DiagnosticEngine::HandlerID, DiagnosticEngine::HandlerTy,
208                        2>
209       handlers;
210 
211   /// This is a unique identifier counter for diagnostic handlers in the
212   /// context. This id starts at 1 to allow for 0 to be used as a sentinel.
213   DiagnosticEngine::HandlerID uniqueHandlerId = 1;
214 };
215 } // namespace detail
216 } // namespace mlir
217 
218 /// Emit a diagnostic using the registered issue handle if present, or with
219 /// the default behavior if not.
220 void DiagnosticEngineImpl::emit(Diagnostic diag) {
221   llvm::sys::SmartScopedLock<true> lock(mutex);
222 
223   // Try to process the given diagnostic on one of the registered handlers.
224   // Handlers are walked in reverse order, so that the most recent handler is
225   // processed first.
226   for (auto &handlerIt : llvm::reverse(handlers))
227     if (succeeded(handlerIt.second(diag)))
228       return;
229 
230   // Otherwise, if this is an error we emit it to stderr.
231   if (diag.getSeverity() != DiagnosticSeverity::Error)
232     return;
233 
234   auto &os = llvm::errs();
235   if (!diag.getLocation().isa<UnknownLoc>())
236     os << diag.getLocation() << ": ";
237   os << "error: ";
238 
239   // The default behavior for errors is to emit them to stderr.
240   os << diag << '\n';
241   os.flush();
242 }
243 
244 //===----------------------------------------------------------------------===//
245 // DiagnosticEngine
246 //===----------------------------------------------------------------------===//
247 
248 DiagnosticEngine::DiagnosticEngine() : impl(new DiagnosticEngineImpl()) {}
249 DiagnosticEngine::~DiagnosticEngine() {}
250 
251 /// Register a new handler for diagnostics to the engine. This function returns
252 /// a unique identifier for the registered handler, which can be used to
253 /// unregister this handler at a later time.
254 auto DiagnosticEngine::registerHandler(const HandlerTy &handler) -> HandlerID {
255   llvm::sys::SmartScopedLock<true> lock(impl->mutex);
256   auto uniqueID = impl->uniqueHandlerId++;
257   impl->handlers.insert({uniqueID, handler});
258   return uniqueID;
259 }
260 
261 /// Erase the registered diagnostic handler with the given identifier.
262 void DiagnosticEngine::eraseHandler(HandlerID handlerID) {
263   llvm::sys::SmartScopedLock<true> lock(impl->mutex);
264   impl->handlers.erase(handlerID);
265 }
266 
267 /// Emit a diagnostic using the registered issue handler if present, or with
268 /// the default behavior if not.
269 void DiagnosticEngine::emit(Diagnostic diag) {
270   assert(diag.getSeverity() != DiagnosticSeverity::Note &&
271          "notes should not be emitted directly");
272   impl->emit(std::move(diag));
273 }
274 
275 /// Helper function used to emit a diagnostic with an optionally empty twine
276 /// message. If the message is empty, then it is not inserted into the
277 /// diagnostic.
278 static InFlightDiagnostic
279 emitDiag(Location location, DiagnosticSeverity severity, const Twine &message) {
280   MLIRContext *ctx = location->getContext();
281   auto &diagEngine = ctx->getDiagEngine();
282   auto diag = diagEngine.emit(location, severity);
283   if (!message.isTriviallyEmpty())
284     diag << message;
285 
286   // Add the stack trace as a note if necessary.
287   if (ctx->shouldPrintStackTraceOnDiagnostic()) {
288     std::string bt;
289     {
290       llvm::raw_string_ostream stream(bt);
291       llvm::sys::PrintStackTrace(stream);
292     }
293     if (!bt.empty())
294       diag.attachNote() << "diagnostic emitted with trace:\n" << bt;
295   }
296 
297   return diag;
298 }
299 
300 /// Emit an error message using this location.
301 InFlightDiagnostic mlir::emitError(Location loc) { return emitError(loc, {}); }
302 InFlightDiagnostic mlir::emitError(Location loc, const Twine &message) {
303   return emitDiag(loc, DiagnosticSeverity::Error, message);
304 }
305 
306 /// Emit a warning message using this location.
307 InFlightDiagnostic mlir::emitWarning(Location loc) {
308   return emitWarning(loc, {});
309 }
310 InFlightDiagnostic mlir::emitWarning(Location loc, const Twine &message) {
311   return emitDiag(loc, DiagnosticSeverity::Warning, message);
312 }
313 
314 /// Emit a remark message using this location.
315 InFlightDiagnostic mlir::emitRemark(Location loc) {
316   return emitRemark(loc, {});
317 }
318 InFlightDiagnostic mlir::emitRemark(Location loc, const Twine &message) {
319   return emitDiag(loc, DiagnosticSeverity::Remark, message);
320 }
321 
322 //===----------------------------------------------------------------------===//
323 // ScopedDiagnosticHandler
324 //===----------------------------------------------------------------------===//
325 
326 ScopedDiagnosticHandler::~ScopedDiagnosticHandler() {
327   if (handlerID)
328     ctx->getDiagEngine().eraseHandler(handlerID);
329 }
330 
331 //===----------------------------------------------------------------------===//
332 // SourceMgrDiagnosticHandler
333 //===----------------------------------------------------------------------===//
334 namespace mlir {
335 namespace detail {
336 struct SourceMgrDiagnosticHandlerImpl {
337   /// Get a memory buffer for the given file, or nullptr if one is not found.
338   const llvm::MemoryBuffer *getBufferForFile(llvm::SourceMgr &mgr,
339                                              StringRef filename) {
340     // Check for an existing mapping to the buffer id for this file.
341     auto bufferIt = filenameToBuf.find(filename);
342     if (bufferIt != filenameToBuf.end())
343       return bufferIt->second;
344 
345     // Look for a buffer in the manager that has this filename.
346     for (unsigned i = 1, e = mgr.getNumBuffers() + 1; i != e; ++i) {
347       auto *buf = mgr.getMemoryBuffer(i);
348       if (buf->getBufferIdentifier() == filename)
349         return filenameToBuf[filename] = buf;
350     }
351 
352     // Otherwise, try to load the source file.
353     const llvm::MemoryBuffer *newBuf = nullptr;
354     std::string ignored;
355     if (auto newBufID =
356             mgr.AddIncludeFile(std::string(filename), llvm::SMLoc(), ignored))
357       newBuf = mgr.getMemoryBuffer(newBufID);
358     return filenameToBuf[filename] = newBuf;
359   }
360 
361   /// Mapping between file name and buffer pointer.
362   llvm::StringMap<const llvm::MemoryBuffer *> filenameToBuf;
363 };
364 } // end namespace detail
365 } // end namespace mlir
366 
367 /// Return a processable FileLineColLoc from the given location.
368 static Optional<FileLineColLoc> getFileLineColLoc(Location loc) {
369   switch (loc->getKind()) {
370   case StandardAttributes::NameLocation:
371     return getFileLineColLoc(loc.cast<NameLoc>().getChildLoc());
372   case StandardAttributes::FileLineColLocation:
373     return loc.cast<FileLineColLoc>();
374   case StandardAttributes::CallSiteLocation:
375     // Process the callee of a callsite location.
376     return getFileLineColLoc(loc.cast<CallSiteLoc>().getCallee());
377   case StandardAttributes::FusedLocation:
378     for (auto subLoc : loc.cast<FusedLoc>().getLocations()) {
379       if (auto callLoc = getFileLineColLoc(subLoc)) {
380         return callLoc;
381       }
382     }
383     return llvm::None;
384   default:
385     return llvm::None;
386   }
387 }
388 
389 /// Return a processable CallSiteLoc from the given location.
390 static Optional<CallSiteLoc> getCallSiteLoc(Location loc) {
391   switch (loc->getKind()) {
392   case StandardAttributes::NameLocation:
393     return getCallSiteLoc(loc.cast<NameLoc>().getChildLoc());
394   case StandardAttributes::CallSiteLocation:
395     return loc.cast<CallSiteLoc>();
396   case StandardAttributes::FusedLocation:
397     for (auto subLoc : loc.cast<FusedLoc>().getLocations()) {
398       if (auto callLoc = getCallSiteLoc(subLoc)) {
399         return callLoc;
400       }
401     }
402     return llvm::None;
403   default:
404     return llvm::None;
405   }
406 }
407 
408 /// Given a diagnostic kind, returns the LLVM DiagKind.
409 static llvm::SourceMgr::DiagKind getDiagKind(DiagnosticSeverity kind) {
410   switch (kind) {
411   case DiagnosticSeverity::Note:
412     return llvm::SourceMgr::DK_Note;
413   case DiagnosticSeverity::Warning:
414     return llvm::SourceMgr::DK_Warning;
415   case DiagnosticSeverity::Error:
416     return llvm::SourceMgr::DK_Error;
417   case DiagnosticSeverity::Remark:
418     return llvm::SourceMgr::DK_Remark;
419   }
420   llvm_unreachable("Unknown DiagnosticSeverity");
421 }
422 
423 SourceMgrDiagnosticHandler::SourceMgrDiagnosticHandler(llvm::SourceMgr &mgr,
424                                                        MLIRContext *ctx,
425                                                        raw_ostream &os)
426     : ScopedDiagnosticHandler(ctx), mgr(mgr), os(os),
427       impl(new SourceMgrDiagnosticHandlerImpl()) {
428   setHandler([this](Diagnostic &diag) { emitDiagnostic(diag); });
429 }
430 
431 SourceMgrDiagnosticHandler::SourceMgrDiagnosticHandler(llvm::SourceMgr &mgr,
432                                                        MLIRContext *ctx)
433     : SourceMgrDiagnosticHandler(mgr, ctx, llvm::errs()) {}
434 
435 SourceMgrDiagnosticHandler::~SourceMgrDiagnosticHandler() {}
436 
437 void SourceMgrDiagnosticHandler::emitDiagnostic(Location loc, Twine message,
438                                                 DiagnosticSeverity kind,
439                                                 bool displaySourceLine) {
440   // Extract a file location from this loc.
441   auto fileLoc = getFileLineColLoc(loc);
442 
443   // If one doesn't exist, then print the raw message without a source location.
444   if (!fileLoc) {
445     std::string str;
446     llvm::raw_string_ostream strOS(str);
447     if (!loc.isa<UnknownLoc>())
448       strOS << loc << ": ";
449     strOS << message;
450     return mgr.PrintMessage(os, llvm::SMLoc(), getDiagKind(kind), strOS.str());
451   }
452 
453   // Otherwise if we are displaying the source line, try to convert the file
454   // location to an SMLoc.
455   if (displaySourceLine) {
456     auto smloc = convertLocToSMLoc(*fileLoc);
457     if (smloc.isValid())
458       return mgr.PrintMessage(os, smloc, getDiagKind(kind), message);
459   }
460 
461   // If the conversion was unsuccessful, create a diagnostic with the file
462   // information. We manually combine the line and column to avoid asserts in
463   // the constructor of SMDiagnostic that takes a location.
464   std::string locStr;
465   llvm::raw_string_ostream locOS(locStr);
466   locOS << fileLoc->getFilename() << ":" << fileLoc->getLine() << ":"
467         << fileLoc->getColumn();
468   llvm::SMDiagnostic diag(locOS.str(), getDiagKind(kind), message.str());
469   diag.print(nullptr, os);
470 }
471 
472 /// Emit the given diagnostic with the held source manager.
473 void SourceMgrDiagnosticHandler::emitDiagnostic(Diagnostic &diag) {
474   // Emit the diagnostic.
475   Location loc = diag.getLocation();
476   emitDiagnostic(loc, diag.str(), diag.getSeverity());
477 
478   // If the diagnostic location was a call site location, then print the call
479   // stack as well.
480   if (auto callLoc = getCallSiteLoc(loc)) {
481     // Print the call stack while valid, or until the limit is reached.
482     loc = callLoc->getCaller();
483     for (unsigned curDepth = 0; curDepth < callStackLimit; ++curDepth) {
484       emitDiagnostic(loc, "called from", DiagnosticSeverity::Note);
485       if ((callLoc = getCallSiteLoc(loc)))
486         loc = callLoc->getCaller();
487       else
488         break;
489     }
490   }
491 
492   // Emit each of the notes. Only display the source code if the location is
493   // different from the previous location.
494   for (auto &note : diag.getNotes()) {
495     emitDiagnostic(note.getLocation(), note.str(), note.getSeverity(),
496                    /*displaySourceLine=*/loc != note.getLocation());
497     loc = note.getLocation();
498   }
499 }
500 
501 /// Get a memory buffer for the given file, or nullptr if one is not found.
502 const llvm::MemoryBuffer *
503 SourceMgrDiagnosticHandler::getBufferForFile(StringRef filename) {
504   return impl->getBufferForFile(mgr, filename);
505 }
506 
507 /// Get a memory buffer for the given file, or the main file of the source
508 /// manager if one doesn't exist. This always returns non-null.
509 llvm::SMLoc SourceMgrDiagnosticHandler::convertLocToSMLoc(FileLineColLoc loc) {
510   // Get the buffer for this filename.
511   auto *membuf = getBufferForFile(loc.getFilename());
512   if (!membuf)
513     return llvm::SMLoc();
514 
515   // TODO: This should really be upstreamed to be a method on llvm::SourceMgr.
516   // Doing so would allow it to use the offset cache that is already maintained
517   // by SrcBuffer, making this more efficient.
518   unsigned lineNo = loc.getLine();
519   unsigned columnNo = loc.getColumn();
520 
521   // Scan for the correct line number.
522   const char *position = membuf->getBufferStart();
523   const char *end = membuf->getBufferEnd();
524 
525   // We start counting line and column numbers from 1.
526   if (lineNo != 0)
527     --lineNo;
528   if (columnNo != 0)
529     --columnNo;
530 
531   while (position < end && lineNo) {
532     auto curChar = *position++;
533 
534     // Scan for newlines.  If this isn't one, ignore it.
535     if (curChar != '\r' && curChar != '\n')
536       continue;
537 
538     // We saw a line break, decrement our counter.
539     --lineNo;
540 
541     // Check for \r\n and \n\r and treat it as a single escape.  We know that
542     // looking past one character is safe because MemoryBuffer's are always nul
543     // terminated.
544     if (*position != curChar && (*position == '\r' || *position == '\n'))
545       ++position;
546   }
547 
548   // If the line/column counter was invalid, return a pointer to the start of
549   // the buffer.
550   if (lineNo || position + columnNo > end)
551     return llvm::SMLoc::getFromPointer(membuf->getBufferStart());
552 
553   // If the column is zero, try to skip to the first non-whitespace character.
554   if (columnNo == 0) {
555     auto isNewline = [](char c) { return c == '\n' || c == '\r'; };
556     auto isWhitespace = [](char c) { return c == ' ' || c == '\t'; };
557 
558     // Look for a valid non-whitespace character before the next line.
559     for (auto *newPos = position; newPos < end && !isNewline(*newPos); ++newPos)
560       if (!isWhitespace(*newPos))
561         return llvm::SMLoc::getFromPointer(newPos);
562   }
563 
564   // Otherwise return the right pointer.
565   return llvm::SMLoc::getFromPointer(position + columnNo);
566 }
567 
568 //===----------------------------------------------------------------------===//
569 // SourceMgrDiagnosticVerifierHandler
570 //===----------------------------------------------------------------------===//
571 
572 namespace mlir {
573 namespace detail {
574 // Record the expected diagnostic's position, substring and whether it was
575 // seen.
576 struct ExpectedDiag {
577   DiagnosticSeverity kind;
578   unsigned lineNo;
579   StringRef substring;
580   llvm::SMLoc fileLoc;
581   bool matched;
582 };
583 
584 struct SourceMgrDiagnosticVerifierHandlerImpl {
585   SourceMgrDiagnosticVerifierHandlerImpl() : status(success()) {}
586 
587   /// Returns the expected diagnostics for the given source file.
588   Optional<MutableArrayRef<ExpectedDiag>> getExpectedDiags(StringRef bufName);
589 
590   /// Computes the expected diagnostics for the given source buffer.
591   MutableArrayRef<ExpectedDiag>
592   computeExpectedDiags(const llvm::MemoryBuffer *buf);
593 
594   /// The current status of the verifier.
595   LogicalResult status;
596 
597   /// A list of expected diagnostics for each buffer of the source manager.
598   llvm::StringMap<SmallVector<ExpectedDiag, 2>> expectedDiagsPerFile;
599 
600   /// Regex to match the expected diagnostics format.
601   llvm::Regex expected = llvm::Regex("expected-(error|note|remark|warning) "
602                                      "*(@([+-][0-9]+|above|below))? *{{(.*)}}");
603 };
604 } // end namespace detail
605 } // end namespace mlir
606 
607 /// Given a diagnostic kind, return a human readable string for it.
608 static StringRef getDiagKindStr(DiagnosticSeverity kind) {
609   switch (kind) {
610   case DiagnosticSeverity::Note:
611     return "note";
612   case DiagnosticSeverity::Warning:
613     return "warning";
614   case DiagnosticSeverity::Error:
615     return "error";
616   case DiagnosticSeverity::Remark:
617     return "remark";
618   }
619   llvm_unreachable("Unknown DiagnosticSeverity");
620 }
621 
622 /// Returns the expected diagnostics for the given source file.
623 Optional<MutableArrayRef<ExpectedDiag>>
624 SourceMgrDiagnosticVerifierHandlerImpl::getExpectedDiags(StringRef bufName) {
625   auto expectedDiags = expectedDiagsPerFile.find(bufName);
626   if (expectedDiags != expectedDiagsPerFile.end())
627     return MutableArrayRef<ExpectedDiag>(expectedDiags->second);
628   return llvm::None;
629 }
630 
631 /// Computes the expected diagnostics for the given source buffer.
632 MutableArrayRef<ExpectedDiag>
633 SourceMgrDiagnosticVerifierHandlerImpl::computeExpectedDiags(
634     const llvm::MemoryBuffer *buf) {
635   // If the buffer is invalid, return an empty list.
636   if (!buf)
637     return llvm::None;
638   auto &expectedDiags = expectedDiagsPerFile[buf->getBufferIdentifier()];
639 
640   // The number of the last line that did not correlate to a designator.
641   unsigned lastNonDesignatorLine = 0;
642 
643   // The indices of designators that apply to the next non designator line.
644   SmallVector<unsigned, 1> designatorsForNextLine;
645 
646   // Scan the file for expected-* designators.
647   SmallVector<StringRef, 100> lines;
648   buf->getBuffer().split(lines, '\n');
649   for (unsigned lineNo = 0, e = lines.size(); lineNo < e; ++lineNo) {
650     SmallVector<StringRef, 4> matches;
651     if (!expected.match(lines[lineNo], &matches)) {
652       // Check for designators that apply to this line.
653       if (!designatorsForNextLine.empty()) {
654         for (unsigned diagIndex : designatorsForNextLine)
655           expectedDiags[diagIndex].lineNo = lineNo + 1;
656         designatorsForNextLine.clear();
657       }
658       lastNonDesignatorLine = lineNo;
659       continue;
660     }
661 
662     // Point to the start of expected-*.
663     auto expectedStart = llvm::SMLoc::getFromPointer(matches[0].data());
664 
665     DiagnosticSeverity kind;
666     if (matches[1] == "error")
667       kind = DiagnosticSeverity::Error;
668     else if (matches[1] == "warning")
669       kind = DiagnosticSeverity::Warning;
670     else if (matches[1] == "remark")
671       kind = DiagnosticSeverity::Remark;
672     else {
673       assert(matches[1] == "note");
674       kind = DiagnosticSeverity::Note;
675     }
676 
677     ExpectedDiag record{kind, lineNo + 1, matches[4], expectedStart, false};
678     auto offsetMatch = matches[2];
679     if (!offsetMatch.empty()) {
680       offsetMatch = offsetMatch.drop_front(1);
681 
682       // Get the integer value without the @ and +/- prefix.
683       if (offsetMatch[0] == '+' || offsetMatch[0] == '-') {
684         int offset;
685         offsetMatch.drop_front().getAsInteger(0, offset);
686 
687         if (offsetMatch.front() == '+')
688           record.lineNo += offset;
689         else
690           record.lineNo -= offset;
691       } else if (offsetMatch.consume_front("above")) {
692         // If the designator applies 'above' we add it to the last non
693         // designator line.
694         record.lineNo = lastNonDesignatorLine + 1;
695       } else {
696         // Otherwise, this is a 'below' designator and applies to the next
697         // non-designator line.
698         assert(offsetMatch.consume_front("below"));
699         designatorsForNextLine.push_back(expectedDiags.size());
700 
701         // Set the line number to the last in the case that this designator ends
702         // up dangling.
703         record.lineNo = e;
704       }
705     }
706     expectedDiags.push_back(record);
707   }
708   return expectedDiags;
709 }
710 
711 SourceMgrDiagnosticVerifierHandler::SourceMgrDiagnosticVerifierHandler(
712     llvm::SourceMgr &srcMgr, MLIRContext *ctx, raw_ostream &out)
713     : SourceMgrDiagnosticHandler(srcMgr, ctx, out),
714       impl(new SourceMgrDiagnosticVerifierHandlerImpl()) {
715   // Compute the expected diagnostics for each of the current files in the
716   // source manager.
717   for (unsigned i = 0, e = mgr.getNumBuffers(); i != e; ++i)
718     (void)impl->computeExpectedDiags(mgr.getMemoryBuffer(i + 1));
719 
720   // Register a handler to verify the diagnostics.
721   setHandler([&](Diagnostic &diag) {
722     // Process the main diagnostics.
723     process(diag);
724 
725     // Process each of the notes.
726     for (auto &note : diag.getNotes())
727       process(note);
728   });
729 }
730 
731 SourceMgrDiagnosticVerifierHandler::SourceMgrDiagnosticVerifierHandler(
732     llvm::SourceMgr &srcMgr, MLIRContext *ctx)
733     : SourceMgrDiagnosticVerifierHandler(srcMgr, ctx, llvm::errs()) {}
734 
735 SourceMgrDiagnosticVerifierHandler::~SourceMgrDiagnosticVerifierHandler() {
736   // Ensure that all expected diagnostics were handled.
737   (void)verify();
738 }
739 
740 /// Returns the status of the verifier and verifies that all expected
741 /// diagnostics were emitted. This return success if all diagnostics were
742 /// verified correctly, failure otherwise.
743 LogicalResult SourceMgrDiagnosticVerifierHandler::verify() {
744   // Verify that all expected errors were seen.
745   for (auto &expectedDiagsPair : impl->expectedDiagsPerFile) {
746     for (auto &err : expectedDiagsPair.second) {
747       if (err.matched)
748         continue;
749       llvm::SMRange range(err.fileLoc,
750                           llvm::SMLoc::getFromPointer(err.fileLoc.getPointer() +
751                                                       err.substring.size()));
752       mgr.PrintMessage(os, err.fileLoc, llvm::SourceMgr::DK_Error,
753                        "expected " + getDiagKindStr(err.kind) + " \"" +
754                            err.substring + "\" was not produced",
755                        range);
756       impl->status = failure();
757     }
758   }
759   impl->expectedDiagsPerFile.clear();
760   return impl->status;
761 }
762 
763 /// Process a single diagnostic.
764 void SourceMgrDiagnosticVerifierHandler::process(Diagnostic &diag) {
765   auto kind = diag.getSeverity();
766 
767   // Process a FileLineColLoc.
768   if (auto fileLoc = getFileLineColLoc(diag.getLocation()))
769     return process(*fileLoc, diag.str(), kind);
770 
771   emitDiagnostic(diag.getLocation(),
772                  "unexpected " + getDiagKindStr(kind) + ": " + diag.str(),
773                  DiagnosticSeverity::Error);
774   impl->status = failure();
775 }
776 
777 /// Process a FileLineColLoc diagnostic.
778 void SourceMgrDiagnosticVerifierHandler::process(FileLineColLoc loc,
779                                                  StringRef msg,
780                                                  DiagnosticSeverity kind) {
781   // Get the expected diagnostics for this file.
782   auto diags = impl->getExpectedDiags(loc.getFilename());
783   if (!diags)
784     diags = impl->computeExpectedDiags(getBufferForFile(loc.getFilename()));
785 
786   // Search for a matching expected diagnostic.
787   // If we find something that is close then emit a more specific error.
788   ExpectedDiag *nearMiss = nullptr;
789 
790   // If this was an expected error, remember that we saw it and return.
791   unsigned line = loc.getLine();
792   for (auto &e : *diags) {
793     if (line == e.lineNo && msg.contains(e.substring)) {
794       if (e.kind == kind) {
795         e.matched = true;
796         return;
797       }
798 
799       // If this only differs based on the diagnostic kind, then consider it
800       // to be a near miss.
801       nearMiss = &e;
802     }
803   }
804 
805   // Otherwise, emit an error for the near miss.
806   if (nearMiss)
807     mgr.PrintMessage(os, nearMiss->fileLoc, llvm::SourceMgr::DK_Error,
808                      "'" + getDiagKindStr(kind) +
809                          "' diagnostic emitted when expecting a '" +
810                          getDiagKindStr(nearMiss->kind) + "'");
811   else
812     emitDiagnostic(loc, "unexpected " + getDiagKindStr(kind) + ": " + msg,
813                    DiagnosticSeverity::Error);
814   impl->status = failure();
815 }
816 
817 //===----------------------------------------------------------------------===//
818 // ParallelDiagnosticHandler
819 //===----------------------------------------------------------------------===//
820 
821 namespace mlir {
822 namespace detail {
823 struct ParallelDiagnosticHandlerImpl : public llvm::PrettyStackTraceEntry {
824   struct ThreadDiagnostic {
825     ThreadDiagnostic(size_t id, Diagnostic diag)
826         : id(id), diag(std::move(diag)) {}
827     bool operator<(const ThreadDiagnostic &rhs) const { return id < rhs.id; }
828 
829     /// The id for this diagnostic, this is used for ordering.
830     /// Note: This id corresponds to the ordered position of the current element
831     ///       being processed by a given thread.
832     size_t id;
833 
834     /// The diagnostic.
835     Diagnostic diag;
836   };
837 
838   ParallelDiagnosticHandlerImpl(MLIRContext *ctx) : handlerID(0), context(ctx) {
839     handlerID = ctx->getDiagEngine().registerHandler([this](Diagnostic &diag) {
840       uint64_t tid = llvm::get_threadid();
841       llvm::sys::SmartScopedLock<true> lock(mutex);
842 
843       // If this thread is not tracked, then return failure to let another
844       // handler process this diagnostic.
845       if (!threadToOrderID.count(tid))
846         return failure();
847 
848       // Append a new diagnostic.
849       diagnostics.emplace_back(threadToOrderID[tid], std::move(diag));
850       return success();
851     });
852   }
853 
854   ~ParallelDiagnosticHandlerImpl() override {
855     // Erase this handler from the context.
856     context->getDiagEngine().eraseHandler(handlerID);
857 
858     // Early exit if there are no diagnostics, this is the common case.
859     if (diagnostics.empty())
860       return;
861 
862     // Emit the diagnostics back to the context.
863     emitDiagnostics([&](Diagnostic diag) {
864       return context->getDiagEngine().emit(std::move(diag));
865     });
866   }
867 
868   /// Utility method to emit any held diagnostics.
869   void emitDiagnostics(std::function<void(Diagnostic)> emitFn) const {
870     // Stable sort all of the diagnostics that were emitted. This creates a
871     // deterministic ordering for the diagnostics based upon which order id they
872     // were emitted for.
873     std::stable_sort(diagnostics.begin(), diagnostics.end());
874 
875     // Emit each diagnostic to the context again.
876     for (ThreadDiagnostic &diag : diagnostics)
877       emitFn(std::move(diag.diag));
878   }
879 
880   /// Set the order id for the current thread.
881   void setOrderIDForThread(size_t orderID) {
882     uint64_t tid = llvm::get_threadid();
883     llvm::sys::SmartScopedLock<true> lock(mutex);
884     threadToOrderID[tid] = orderID;
885   }
886 
887   /// Remove the order id for the current thread.
888   void eraseOrderIDForThread() {
889     uint64_t tid = llvm::get_threadid();
890     llvm::sys::SmartScopedLock<true> lock(mutex);
891     threadToOrderID.erase(tid);
892   }
893 
894   /// Dump the current diagnostics that were inflight.
895   void print(raw_ostream &os) const override {
896     // Early exit if there are no diagnostics, this is the common case.
897     if (diagnostics.empty())
898       return;
899 
900     os << "In-Flight Diagnostics:\n";
901     emitDiagnostics([&](Diagnostic diag) {
902       os.indent(4);
903 
904       // Print each diagnostic with the format:
905       //   "<location>: <kind>: <msg>"
906       if (!diag.getLocation().isa<UnknownLoc>())
907         os << diag.getLocation() << ": ";
908       switch (diag.getSeverity()) {
909       case DiagnosticSeverity::Error:
910         os << "error: ";
911         break;
912       case DiagnosticSeverity::Warning:
913         os << "warning: ";
914         break;
915       case DiagnosticSeverity::Note:
916         os << "note: ";
917         break;
918       case DiagnosticSeverity::Remark:
919         os << "remark: ";
920         break;
921       }
922       os << diag << '\n';
923     });
924   }
925 
926   /// A smart mutex to lock access to the internal state.
927   llvm::sys::SmartMutex<true> mutex;
928 
929   /// A mapping between the thread id and the current order id.
930   DenseMap<uint64_t, size_t> threadToOrderID;
931 
932   /// An unordered list of diagnostics that were emitted.
933   mutable std::vector<ThreadDiagnostic> diagnostics;
934 
935   /// The unique id for the parallel handler.
936   DiagnosticEngine::HandlerID handlerID;
937 
938   /// The context to emit the diagnostics to.
939   MLIRContext *context;
940 };
941 } // end namespace detail
942 } // end namespace mlir
943 
944 ParallelDiagnosticHandler::ParallelDiagnosticHandler(MLIRContext *ctx)
945     : impl(new ParallelDiagnosticHandlerImpl(ctx)) {}
946 ParallelDiagnosticHandler::~ParallelDiagnosticHandler() {}
947 
948 /// Set the order id for the current thread.
949 void ParallelDiagnosticHandler::setOrderIDForThread(size_t orderID) {
950   impl->setOrderIDForThread(orderID);
951 }
952 
953 /// Remove the order id for the current thread. This removes the thread from
954 /// diagnostics tracking.
955 void ParallelDiagnosticHandler::eraseOrderIDForThread() {
956   impl->eraseOrderIDForThread();
957 }
958