1 //===- MLIRServer.cpp - MLIR Generic Language Server ----------------------===//
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 "MLIRServer.h"
10 #include "lsp/Logging.h"
11 #include "lsp/Protocol.h"
12 #include "mlir/IR/Operation.h"
13 #include "mlir/Parser.h"
14 #include "mlir/Parser/AsmParserState.h"
15 #include "llvm/Support/SourceMgr.h"
16 
17 using namespace mlir;
18 
19 /// Returns a language server position for the given source location.
20 static lsp::Position getPosFromLoc(llvm::SourceMgr &mgr, llvm::SMLoc loc) {
21   std::pair<unsigned, unsigned> lineAndCol = mgr.getLineAndColumn(loc);
22   lsp::Position pos;
23   pos.line = lineAndCol.first - 1;
24   pos.character = lineAndCol.second;
25   return pos;
26 }
27 
28 /// Returns a source location from the given language server position.
29 static llvm::SMLoc getPosFromLoc(llvm::SourceMgr &mgr, lsp::Position pos) {
30   return mgr.FindLocForLineAndColumn(mgr.getMainFileID(), pos.line + 1,
31                                      pos.character);
32 }
33 
34 /// Returns a language server range for the given source range.
35 static lsp::Range getRangeFromLoc(llvm::SourceMgr &mgr, llvm::SMRange range) {
36   // lsp::Range is an inclusive range, SMRange is half-open.
37   llvm::SMLoc inclusiveEnd =
38       llvm::SMLoc::getFromPointer(range.End.getPointer() - 1);
39   return {getPosFromLoc(mgr, range.Start), getPosFromLoc(mgr, inclusiveEnd)};
40 }
41 
42 /// Returns a language server location from the given source range.
43 static lsp::Location getLocationFromLoc(llvm::SourceMgr &mgr,
44                                         llvm::SMRange range,
45                                         const lsp::URIForFile &uri) {
46   return lsp::Location{uri, getRangeFromLoc(mgr, range)};
47 }
48 
49 /// Returns a language server location from the given MLIR file location.
50 static Optional<lsp::Location> getLocationFromLoc(FileLineColLoc loc) {
51   llvm::Expected<lsp::URIForFile> sourceURI =
52       lsp::URIForFile::fromFile(loc.getFilename());
53   if (!sourceURI) {
54     lsp::Logger::error("Failed to create URI for file `{0}`: {1}",
55                        loc.getFilename(),
56                        llvm::toString(sourceURI.takeError()));
57     return llvm::None;
58   }
59 
60   lsp::Position position;
61   position.line = loc.getLine() - 1;
62   position.character = loc.getColumn();
63   return lsp::Location{*sourceURI, lsp::Range(position)};
64 }
65 
66 /// Returns a language server location from the given MLIR location, or None if
67 /// one couldn't be created. `uri` is an optional additional filter that, when
68 /// present, is used to filter sub locations that do not share the same uri.
69 static Optional<lsp::Location>
70 getLocationFromLoc(llvm::SourceMgr &sourceMgr, Location loc,
71                    const lsp::URIForFile *uri = nullptr) {
72   Optional<lsp::Location> location;
73   loc->walk([&](Location nestedLoc) {
74     FileLineColLoc fileLoc = nestedLoc.dyn_cast<FileLineColLoc>();
75     if (!fileLoc)
76       return WalkResult::advance();
77 
78     Optional<lsp::Location> sourceLoc = getLocationFromLoc(fileLoc);
79     if (sourceLoc && (!uri || sourceLoc->uri == *uri)) {
80       location = *sourceLoc;
81       llvm::SMLoc loc = sourceMgr.FindLocForLineAndColumn(
82           sourceMgr.getMainFileID(), fileLoc.getLine(), fileLoc.getColumn());
83 
84       // Use range of potential identifier starting at location, else length 1
85       // range.
86       location->range.end.character += 1;
87       if (Optional<llvm::SMRange> range =
88               AsmParserState::convertIdLocToRange(loc)) {
89         auto lineCol = sourceMgr.getLineAndColumn(range->End);
90         location->range.end.character =
91             std::max(fileLoc.getColumn() + 1, lineCol.second - 1);
92       }
93       return WalkResult::interrupt();
94     }
95     return WalkResult::advance();
96   });
97   return location;
98 }
99 
100 /// Collect all of the locations from the given MLIR location that are not
101 /// contained within the given URI.
102 static void collectLocationsFromLoc(Location loc,
103                                     std::vector<lsp::Location> &locations,
104                                     const lsp::URIForFile &uri) {
105   SetVector<Location> visitedLocs;
106   loc->walk([&](Location nestedLoc) {
107     FileLineColLoc fileLoc = nestedLoc.dyn_cast<FileLineColLoc>();
108     if (!fileLoc || !visitedLocs.insert(nestedLoc))
109       return WalkResult::advance();
110 
111     Optional<lsp::Location> sourceLoc = getLocationFromLoc(fileLoc);
112     if (sourceLoc && sourceLoc->uri != uri)
113       locations.push_back(*sourceLoc);
114     return WalkResult::advance();
115   });
116 }
117 
118 /// Returns true if the given range contains the given source location. Note
119 /// that this has slightly different behavior than SMRange because it is
120 /// inclusive of the end location.
121 static bool contains(llvm::SMRange range, llvm::SMLoc loc) {
122   return range.Start.getPointer() <= loc.getPointer() &&
123          loc.getPointer() <= range.End.getPointer();
124 }
125 
126 /// Returns true if the given location is contained by the definition or one of
127 /// the uses of the given SMDefinition. If provided, `overlappedRange` is set to
128 /// the range within `def` that the provided `loc` overlapped with.
129 static bool isDefOrUse(const AsmParserState::SMDefinition &def, llvm::SMLoc loc,
130                        llvm::SMRange *overlappedRange = nullptr) {
131   // Check the main definition.
132   if (contains(def.loc, loc)) {
133     if (overlappedRange)
134       *overlappedRange = def.loc;
135     return true;
136   }
137 
138   // Check the uses.
139   auto useIt = llvm::find_if(def.uses, [&](const llvm::SMRange &range) {
140     return contains(range, loc);
141   });
142   if (useIt != def.uses.end()) {
143     if (overlappedRange)
144       *overlappedRange = *useIt;
145     return true;
146   }
147   return false;
148 }
149 
150 /// Given a location pointing to a result, return the result number it refers
151 /// to or None if it refers to all of the results.
152 static Optional<unsigned> getResultNumberFromLoc(llvm::SMLoc loc) {
153   // Skip all of the identifier characters.
154   auto isIdentifierChar = [](char c) {
155     return isalnum(c) || c == '%' || c == '$' || c == '.' || c == '_' ||
156            c == '-';
157   };
158   const char *curPtr = loc.getPointer();
159   while (isIdentifierChar(*curPtr))
160     ++curPtr;
161 
162   // Check to see if this location indexes into the result group, via `#`. If it
163   // doesn't, we can't extract a sub result number.
164   if (*curPtr != '#')
165     return llvm::None;
166 
167   // Compute the sub result number from the remaining portion of the string.
168   const char *numberStart = ++curPtr;
169   while (llvm::isDigit(*curPtr))
170     ++curPtr;
171   StringRef numberStr(numberStart, curPtr - numberStart);
172   unsigned resultNumber = 0;
173   return numberStr.consumeInteger(10, resultNumber) ? Optional<unsigned>()
174                                                     : resultNumber;
175 }
176 
177 /// Given a source location range, return the text covered by the given range.
178 /// If the range is invalid, returns None.
179 static Optional<StringRef> getTextFromRange(llvm::SMRange range) {
180   if (!range.isValid())
181     return None;
182   const char *startPtr = range.Start.getPointer();
183   return StringRef(startPtr, range.End.getPointer() - startPtr);
184 }
185 
186 /// Given a block, return its position in its parent region.
187 static unsigned getBlockNumber(Block *block) {
188   return std::distance(block->getParent()->begin(), block->getIterator());
189 }
190 
191 /// Given a block and source location, print the source name of the block to the
192 /// given output stream.
193 static void printDefBlockName(raw_ostream &os, Block *block,
194                               llvm::SMRange loc = {}) {
195   // Try to extract a name from the source location.
196   Optional<StringRef> text = getTextFromRange(loc);
197   if (text && text->startswith("^")) {
198     os << *text;
199     return;
200   }
201 
202   // Otherwise, we don't have a name so print the block number.
203   os << "<Block #" << getBlockNumber(block) << ">";
204 }
205 static void printDefBlockName(raw_ostream &os,
206                               const AsmParserState::BlockDefinition &def) {
207   printDefBlockName(os, def.block, def.definition.loc);
208 }
209 
210 /// Convert the given MLIR diagnostic to the LSP form.
211 static lsp::Diagnostic getLspDiagnoticFromDiag(llvm::SourceMgr &sourceMgr,
212                                                Diagnostic &diag,
213                                                const lsp::URIForFile &uri) {
214   lsp::Diagnostic lspDiag;
215   lspDiag.source = "mlir";
216 
217   // Note: Right now all of the diagnostics are treated as parser issues, but
218   // some are parser and some are verifier.
219   lspDiag.category = "Parse Error";
220 
221   // Try to grab a file location for this diagnostic.
222   // TODO: For simplicity, we just grab the first one. It may be likely that we
223   // will need a more interesting heuristic here.'
224   Optional<lsp::Location> lspLocation =
225       getLocationFromLoc(sourceMgr, diag.getLocation(), &uri);
226   if (lspLocation)
227     lspDiag.range = lspLocation->range;
228 
229   // Convert the severity for the diagnostic.
230   switch (diag.getSeverity()) {
231   case DiagnosticSeverity::Note:
232     llvm_unreachable("expected notes to be handled separately");
233   case DiagnosticSeverity::Warning:
234     lspDiag.severity = lsp::DiagnosticSeverity::Warning;
235     break;
236   case DiagnosticSeverity::Error:
237     lspDiag.severity = lsp::DiagnosticSeverity::Error;
238     break;
239   case DiagnosticSeverity::Remark:
240     lspDiag.severity = lsp::DiagnosticSeverity::Information;
241     break;
242   }
243   lspDiag.message = diag.str();
244 
245   // Attach any notes to the main diagnostic as related information.
246   std::vector<lsp::DiagnosticRelatedInformation> relatedDiags;
247   for (Diagnostic &note : diag.getNotes()) {
248     lsp::Location noteLoc;
249     if (Optional<lsp::Location> loc =
250             getLocationFromLoc(sourceMgr, note.getLocation()))
251       noteLoc = *loc;
252     else
253       noteLoc.uri = uri;
254     relatedDiags.emplace_back(noteLoc, note.str());
255   }
256   if (!relatedDiags.empty())
257     lspDiag.relatedInformation = std::move(relatedDiags);
258 
259   return lspDiag;
260 }
261 
262 //===----------------------------------------------------------------------===//
263 // MLIRDocument
264 //===----------------------------------------------------------------------===//
265 
266 namespace {
267 /// This class represents all of the information pertaining to a specific MLIR
268 /// document.
269 struct MLIRDocument {
270   MLIRDocument(const lsp::URIForFile &uri, StringRef contents,
271                DialectRegistry &registry,
272                std::vector<lsp::Diagnostic> &diagnostics);
273   MLIRDocument(const MLIRDocument &) = delete;
274   MLIRDocument &operator=(const MLIRDocument &) = delete;
275 
276   //===--------------------------------------------------------------------===//
277   // Definitions and References
278   //===--------------------------------------------------------------------===//
279 
280   void getLocationsOf(const lsp::URIForFile &uri, const lsp::Position &defPos,
281                       std::vector<lsp::Location> &locations);
282   void findReferencesOf(const lsp::URIForFile &uri, const lsp::Position &pos,
283                         std::vector<lsp::Location> &references);
284 
285   //===--------------------------------------------------------------------===//
286   // Hover
287   //===--------------------------------------------------------------------===//
288 
289   Optional<lsp::Hover> findHover(const lsp::URIForFile &uri,
290                                  const lsp::Position &hoverPos);
291   Optional<lsp::Hover>
292   buildHoverForOperation(const AsmParserState::OperationDefinition &op);
293   lsp::Hover buildHoverForOperationResult(llvm::SMRange hoverRange,
294                                           Operation *op, unsigned resultStart,
295                                           unsigned resultEnd,
296                                           llvm::SMLoc posLoc);
297   lsp::Hover buildHoverForBlock(llvm::SMRange hoverRange,
298                                 const AsmParserState::BlockDefinition &block);
299   lsp::Hover
300   buildHoverForBlockArgument(llvm::SMRange hoverRange, BlockArgument arg,
301                              const AsmParserState::BlockDefinition &block);
302 
303   /// The context used to hold the state contained by the parsed document.
304   MLIRContext context;
305 
306   /// The high level parser state used to find definitions and references within
307   /// the source file.
308   AsmParserState asmState;
309 
310   /// The container for the IR parsed from the input file.
311   Block parsedIR;
312 
313   /// The source manager containing the contents of the input file.
314   llvm::SourceMgr sourceMgr;
315 };
316 } // namespace
317 
318 MLIRDocument::MLIRDocument(const lsp::URIForFile &uri, StringRef contents,
319                            DialectRegistry &registry,
320                            std::vector<lsp::Diagnostic> &diagnostics)
321     : context(registry) {
322   context.allowUnregisteredDialects();
323   ScopedDiagnosticHandler handler(&context, [&](Diagnostic &diag) {
324     diagnostics.push_back(getLspDiagnoticFromDiag(sourceMgr, diag, uri));
325   });
326 
327   // Try to parsed the given IR string.
328   auto memBuffer = llvm::MemoryBuffer::getMemBufferCopy(contents, uri.file());
329   if (!memBuffer) {
330     lsp::Logger::error("Failed to create memory buffer for file", uri.file());
331     return;
332   }
333 
334   sourceMgr.AddNewSourceBuffer(std::move(memBuffer), llvm::SMLoc());
335   if (failed(parseSourceFile(sourceMgr, &parsedIR, &context, nullptr,
336                              &asmState))) {
337     // If parsing failed, clear out any of the current state.
338     parsedIR.clear();
339     asmState = AsmParserState();
340     return;
341   }
342 }
343 
344 //===----------------------------------------------------------------------===//
345 // MLIRDocument: Definitions and References
346 //===----------------------------------------------------------------------===//
347 
348 void MLIRDocument::getLocationsOf(const lsp::URIForFile &uri,
349                                   const lsp::Position &defPos,
350                                   std::vector<lsp::Location> &locations) {
351   llvm::SMLoc posLoc = getPosFromLoc(sourceMgr, defPos);
352 
353   // Functor used to check if an SM definition contains the position.
354   auto containsPosition = [&](const AsmParserState::SMDefinition &def) {
355     if (!isDefOrUse(def, posLoc))
356       return false;
357     locations.push_back(getLocationFromLoc(sourceMgr, def.loc, uri));
358     return true;
359   };
360 
361   // Check all definitions related to operations.
362   for (const AsmParserState::OperationDefinition &op : asmState.getOpDefs()) {
363     if (contains(op.loc, posLoc))
364       return collectLocationsFromLoc(op.op->getLoc(), locations, uri);
365     for (const auto &result : op.resultGroups)
366       if (containsPosition(result.second))
367         return collectLocationsFromLoc(op.op->getLoc(), locations, uri);
368   }
369 
370   // Check all definitions related to blocks.
371   for (const AsmParserState::BlockDefinition &block : asmState.getBlockDefs()) {
372     if (containsPosition(block.definition))
373       return;
374     for (const AsmParserState::SMDefinition &arg : block.arguments)
375       if (containsPosition(arg))
376         return;
377   }
378 }
379 
380 void MLIRDocument::findReferencesOf(const lsp::URIForFile &uri,
381                                     const lsp::Position &pos,
382                                     std::vector<lsp::Location> &references) {
383   // Functor used to append all of the definitions/uses of the given SM
384   // definition to the reference list.
385   auto appendSMDef = [&](const AsmParserState::SMDefinition &def) {
386     references.push_back(getLocationFromLoc(sourceMgr, def.loc, uri));
387     for (const llvm::SMRange &use : def.uses)
388       references.push_back(getLocationFromLoc(sourceMgr, use, uri));
389   };
390 
391   llvm::SMLoc posLoc = getPosFromLoc(sourceMgr, pos);
392 
393   // Check all definitions related to operations.
394   for (const AsmParserState::OperationDefinition &op : asmState.getOpDefs()) {
395     if (contains(op.loc, posLoc)) {
396       for (const auto &result : op.resultGroups)
397         appendSMDef(result.second);
398       return;
399     }
400     for (const auto &result : op.resultGroups)
401       if (isDefOrUse(result.second, posLoc))
402         return appendSMDef(result.second);
403   }
404 
405   // Check all definitions related to blocks.
406   for (const AsmParserState::BlockDefinition &block : asmState.getBlockDefs()) {
407     if (isDefOrUse(block.definition, posLoc))
408       return appendSMDef(block.definition);
409 
410     for (const AsmParserState::SMDefinition &arg : block.arguments)
411       if (isDefOrUse(arg, posLoc))
412         return appendSMDef(arg);
413   }
414 }
415 
416 //===----------------------------------------------------------------------===//
417 // MLIRDocument: Hover
418 //===----------------------------------------------------------------------===//
419 
420 Optional<lsp::Hover> MLIRDocument::findHover(const lsp::URIForFile &uri,
421                                              const lsp::Position &hoverPos) {
422   llvm::SMLoc posLoc = getPosFromLoc(sourceMgr, hoverPos);
423   llvm::SMRange hoverRange;
424 
425   // Check for Hovers on operations and results.
426   for (const AsmParserState::OperationDefinition &op : asmState.getOpDefs()) {
427     // Check if the position points at this operation.
428     if (contains(op.loc, posLoc))
429       return buildHoverForOperation(op);
430 
431     // Check if the position points at a result group.
432     for (unsigned i = 0, e = op.resultGroups.size(); i < e; ++i) {
433       const auto &result = op.resultGroups[i];
434       if (!isDefOrUse(result.second, posLoc, &hoverRange))
435         continue;
436 
437       // Get the range of results covered by the over position.
438       unsigned resultStart = result.first;
439       unsigned resultEnd =
440           (i == e - 1) ? op.op->getNumResults() : op.resultGroups[i + 1].first;
441       return buildHoverForOperationResult(hoverRange, op.op, resultStart,
442                                           resultEnd, posLoc);
443     }
444   }
445 
446   // Check to see if the hover is over a block argument.
447   for (const AsmParserState::BlockDefinition &block : asmState.getBlockDefs()) {
448     if (isDefOrUse(block.definition, posLoc, &hoverRange))
449       return buildHoverForBlock(hoverRange, block);
450 
451     for (const auto &arg : llvm::enumerate(block.arguments)) {
452       if (!isDefOrUse(arg.value(), posLoc, &hoverRange))
453         continue;
454 
455       return buildHoverForBlockArgument(
456           hoverRange, block.block->getArgument(arg.index()), block);
457     }
458   }
459   return llvm::None;
460 }
461 
462 Optional<lsp::Hover> MLIRDocument::buildHoverForOperation(
463     const AsmParserState::OperationDefinition &op) {
464   // Don't show hovers for operations with regions to avoid huge hover  blocks.
465   // TODO: Should we add support for printing an op without its regions?
466   if (llvm::any_of(op.op->getRegions(),
467                    [](Region &region) { return !region.empty(); }))
468     return llvm::None;
469 
470   lsp::Hover hover(getRangeFromLoc(sourceMgr, op.loc));
471   llvm::raw_string_ostream os(hover.contents.value);
472 
473   // For hovers on an operation, show the generic form.
474   os << "```mlir\n";
475   op.op->print(
476       os, OpPrintingFlags().printGenericOpForm().elideLargeElementsAttrs());
477   os << "\n```\n";
478 
479   return hover;
480 }
481 
482 lsp::Hover MLIRDocument::buildHoverForOperationResult(llvm::SMRange hoverRange,
483                                                       Operation *op,
484                                                       unsigned resultStart,
485                                                       unsigned resultEnd,
486                                                       llvm::SMLoc posLoc) {
487   lsp::Hover hover(getRangeFromLoc(sourceMgr, hoverRange));
488   llvm::raw_string_ostream os(hover.contents.value);
489 
490   // Add the parent operation name to the hover.
491   os << "Operation: \"" << op->getName() << "\"\n\n";
492 
493   // Check to see if the location points to a specific result within the
494   // group.
495   if (Optional<unsigned> resultNumber = getResultNumberFromLoc(posLoc)) {
496     if ((resultStart + *resultNumber) < resultEnd) {
497       resultStart += *resultNumber;
498       resultEnd = resultStart + 1;
499     }
500   }
501 
502   // Add the range of results and their types to the hover info.
503   if ((resultStart + 1) == resultEnd) {
504     os << "Result #" << resultStart << "\n\n"
505        << "Type: `" << op->getResult(resultStart).getType() << "`\n\n";
506   } else {
507     os << "Result #[" << resultStart << ", " << (resultEnd - 1) << "]\n\n"
508        << "Types: ";
509     llvm::interleaveComma(
510         op->getResults().slice(resultStart, resultEnd), os,
511         [&](Value result) { os << "`" << result.getType() << "`"; });
512   }
513 
514   return hover;
515 }
516 
517 lsp::Hover
518 MLIRDocument::buildHoverForBlock(llvm::SMRange hoverRange,
519                                  const AsmParserState::BlockDefinition &block) {
520   lsp::Hover hover(getRangeFromLoc(sourceMgr, hoverRange));
521   llvm::raw_string_ostream os(hover.contents.value);
522 
523   // Print the given block to the hover output stream.
524   auto printBlockToHover = [&](Block *newBlock) {
525     if (const auto *def = asmState.getBlockDef(newBlock))
526       printDefBlockName(os, *def);
527     else
528       printDefBlockName(os, newBlock);
529   };
530 
531   // Display the parent operation, block number, predecessors, and successors.
532   os << "Operation: \"" << block.block->getParentOp()->getName() << "\"\n\n"
533      << "Block #" << getBlockNumber(block.block) << "\n\n";
534   if (!block.block->hasNoPredecessors()) {
535     os << "Predecessors: ";
536     llvm::interleaveComma(block.block->getPredecessors(), os,
537                           printBlockToHover);
538     os << "\n\n";
539   }
540   if (!block.block->hasNoSuccessors()) {
541     os << "Successors: ";
542     llvm::interleaveComma(block.block->getSuccessors(), os, printBlockToHover);
543     os << "\n\n";
544   }
545 
546   return hover;
547 }
548 
549 lsp::Hover MLIRDocument::buildHoverForBlockArgument(
550     llvm::SMRange hoverRange, BlockArgument arg,
551     const AsmParserState::BlockDefinition &block) {
552   lsp::Hover hover(getRangeFromLoc(sourceMgr, hoverRange));
553   llvm::raw_string_ostream os(hover.contents.value);
554 
555   // Display the parent operation, block, the argument number, and the type.
556   os << "Operation: \"" << block.block->getParentOp()->getName() << "\"\n\n"
557      << "Block: ";
558   printDefBlockName(os, block);
559   os << "\n\nArgument #" << arg.getArgNumber() << "\n\n"
560      << "Type: `" << arg.getType() << "`\n\n";
561 
562   return hover;
563 }
564 
565 //===----------------------------------------------------------------------===//
566 // MLIRTextFileChunk
567 //===----------------------------------------------------------------------===//
568 
569 namespace {
570 /// This class represents a single chunk of an MLIR text file.
571 struct MLIRTextFileChunk {
572   MLIRTextFileChunk(uint64_t lineOffset, const lsp::URIForFile &uri,
573                     StringRef contents, DialectRegistry &registry,
574                     std::vector<lsp::Diagnostic> &diagnostics)
575       : lineOffset(lineOffset), document(uri, contents, registry, diagnostics) {
576   }
577 
578   /// Adjust the line number of the given range to anchor at the beginning of
579   /// the file, instead of the beginning of this chunk.
580   void adjustLocForChunkOffset(lsp::Range &range) {
581     adjustLocForChunkOffset(range.start);
582     adjustLocForChunkOffset(range.end);
583   }
584   /// Adjust the line number of the given position to anchor at the beginning of
585   /// the file, instead of the beginning of this chunk.
586   void adjustLocForChunkOffset(lsp::Position &pos) { pos.line += lineOffset; }
587 
588   /// The line offset of this chunk from the beginning of the file.
589   uint64_t lineOffset;
590   /// The document referred to by this chunk.
591   MLIRDocument document;
592 };
593 } // namespace
594 
595 //===----------------------------------------------------------------------===//
596 // MLIRTextFile
597 //===----------------------------------------------------------------------===//
598 
599 namespace {
600 /// This class represents a text file containing one or more MLIR documents.
601 class MLIRTextFile {
602 public:
603   MLIRTextFile(const lsp::URIForFile &uri, StringRef fileContents,
604                int64_t version, DialectRegistry &registry,
605                std::vector<lsp::Diagnostic> &diagnostics);
606 
607   /// Return the current version of this text file.
608   int64_t getVersion() const { return version; }
609 
610   //===--------------------------------------------------------------------===//
611   // LSP Queries
612   //===--------------------------------------------------------------------===//
613 
614   void getLocationsOf(const lsp::URIForFile &uri, lsp::Position defPos,
615                       std::vector<lsp::Location> &locations);
616   void findReferencesOf(const lsp::URIForFile &uri, lsp::Position pos,
617                         std::vector<lsp::Location> &references);
618   Optional<lsp::Hover> findHover(const lsp::URIForFile &uri,
619                                  lsp::Position hoverPos);
620 
621 private:
622   /// Find the MLIR document that contains the given position, and update the
623   /// position to be anchored at the start of the found chunk instead of the
624   /// beginning of the file.
625   MLIRTextFileChunk &getChunkFor(lsp::Position &pos);
626 
627   /// The full string contents of the file.
628   std::string contents;
629 
630   /// The version of this file.
631   int64_t version;
632 
633   /// The chunks of this file. The order of these chunks is the order in which
634   /// they appear in the text file.
635   std::vector<std::unique_ptr<MLIRTextFileChunk>> chunks;
636 };
637 } // namespace
638 
639 MLIRTextFile::MLIRTextFile(const lsp::URIForFile &uri, StringRef fileContents,
640                            int64_t version, DialectRegistry &registry,
641                            std::vector<lsp::Diagnostic> &diagnostics)
642     : contents(fileContents.str()), version(version) {
643   // Split the file into separate MLIR documents.
644   // TODO: Find a way to share the split file marker with other tools. We don't
645   // want to use `splitAndProcessBuffer` here, but we do want to make sure this
646   // marker doesn't go out of sync.
647   SmallVector<StringRef, 8> subContents;
648   StringRef(contents).split(subContents, "// -----");
649   chunks.emplace_back(std::make_unique<MLIRTextFileChunk>(
650       /*lineOffset=*/0, uri, subContents.front(), registry, diagnostics));
651 
652   uint64_t lineOffset = subContents.front().count('\n');
653   for (StringRef docContents : llvm::drop_begin(subContents)) {
654     unsigned currentNumDiags = diagnostics.size();
655     auto chunk = std::make_unique<MLIRTextFileChunk>(
656         lineOffset, uri, docContents, registry, diagnostics);
657     lineOffset += docContents.count('\n');
658 
659     // Adjust locations used in diagnostics to account for the offset from the
660     // beginning of the file.
661     for (lsp::Diagnostic &diag :
662          llvm::drop_begin(diagnostics, currentNumDiags)) {
663       chunk->adjustLocForChunkOffset(diag.range);
664 
665       if (!diag.relatedInformation)
666         continue;
667       for (auto &it : *diag.relatedInformation)
668         if (it.location.uri == uri)
669           chunk->adjustLocForChunkOffset(it.location.range);
670     }
671     chunks.emplace_back(std::move(chunk));
672   }
673 }
674 
675 void MLIRTextFile::getLocationsOf(const lsp::URIForFile &uri,
676                                   lsp::Position defPos,
677                                   std::vector<lsp::Location> &locations) {
678   MLIRTextFileChunk &chunk = getChunkFor(defPos);
679   chunk.document.getLocationsOf(uri, defPos, locations);
680 
681   // Adjust any locations within this file for the offset of this chunk.
682   if (chunk.lineOffset == 0)
683     return;
684   for (lsp::Location &loc : locations)
685     if (loc.uri == uri)
686       chunk.adjustLocForChunkOffset(loc.range);
687 }
688 
689 void MLIRTextFile::findReferencesOf(const lsp::URIForFile &uri,
690                                     lsp::Position pos,
691                                     std::vector<lsp::Location> &references) {
692   MLIRTextFileChunk &chunk = getChunkFor(pos);
693   chunk.document.findReferencesOf(uri, pos, references);
694 
695   // Adjust any locations within this file for the offset of this chunk.
696   if (chunk.lineOffset == 0)
697     return;
698   for (lsp::Location &loc : references)
699     if (loc.uri == uri)
700       chunk.adjustLocForChunkOffset(loc.range);
701 }
702 
703 Optional<lsp::Hover> MLIRTextFile::findHover(const lsp::URIForFile &uri,
704                                              lsp::Position hoverPos) {
705   MLIRTextFileChunk &chunk = getChunkFor(hoverPos);
706   Optional<lsp::Hover> hoverInfo = chunk.document.findHover(uri, hoverPos);
707 
708   // Adjust any locations within this file for the offset of this chunk.
709   if (chunk.lineOffset != 0 && hoverInfo && hoverInfo->range)
710     chunk.adjustLocForChunkOffset(*hoverInfo->range);
711   return hoverInfo;
712 }
713 
714 MLIRTextFileChunk &MLIRTextFile::getChunkFor(lsp::Position &pos) {
715   if (chunks.size() == 1)
716     return *chunks.front();
717 
718   // Search for the first chunk with a greater line offset, the previous chunk
719   // is the one that contains `pos`.
720   auto it = llvm::upper_bound(
721       chunks, pos, [](const lsp::Position &pos, const auto &chunk) {
722         return static_cast<uint64_t>(pos.line) < chunk->lineOffset;
723       });
724   MLIRTextFileChunk &chunk = it == chunks.end() ? *chunks.back() : **(--it);
725   pos.line -= chunk.lineOffset;
726   return chunk;
727 }
728 
729 //===----------------------------------------------------------------------===//
730 // MLIRServer::Impl
731 //===----------------------------------------------------------------------===//
732 
733 struct lsp::MLIRServer::Impl {
734   Impl(DialectRegistry &registry) : registry(registry) {}
735 
736   /// The registry containing dialects that can be recognized in parsed .mlir
737   /// files.
738   DialectRegistry &registry;
739 
740   /// The files held by the server, mapped by their URI file name.
741   llvm::StringMap<std::unique_ptr<MLIRTextFile>> files;
742 };
743 
744 //===----------------------------------------------------------------------===//
745 // MLIRServer
746 //===----------------------------------------------------------------------===//
747 
748 lsp::MLIRServer::MLIRServer(DialectRegistry &registry)
749     : impl(std::make_unique<Impl>(registry)) {}
750 lsp::MLIRServer::~MLIRServer() {}
751 
752 void lsp::MLIRServer::addOrUpdateDocument(
753     const URIForFile &uri, StringRef contents, int64_t version,
754     std::vector<Diagnostic> &diagnostics) {
755   impl->files[uri.file()] = std::make_unique<MLIRTextFile>(
756       uri, contents, version, impl->registry, diagnostics);
757 }
758 
759 Optional<int64_t> lsp::MLIRServer::removeDocument(const URIForFile &uri) {
760   auto it = impl->files.find(uri.file());
761   if (it == impl->files.end())
762     return llvm::None;
763 
764   int64_t version = it->second->getVersion();
765   impl->files.erase(it);
766   return version;
767 }
768 
769 void lsp::MLIRServer::getLocationsOf(const URIForFile &uri,
770                                      const Position &defPos,
771                                      std::vector<Location> &locations) {
772   auto fileIt = impl->files.find(uri.file());
773   if (fileIt != impl->files.end())
774     fileIt->second->getLocationsOf(uri, defPos, locations);
775 }
776 
777 void lsp::MLIRServer::findReferencesOf(const URIForFile &uri,
778                                        const Position &pos,
779                                        std::vector<Location> &references) {
780   auto fileIt = impl->files.find(uri.file());
781   if (fileIt != impl->files.end())
782     fileIt->second->findReferencesOf(uri, pos, references);
783 }
784 
785 Optional<lsp::Hover> lsp::MLIRServer::findHover(const URIForFile &uri,
786                                                 const Position &hoverPos) {
787   auto fileIt = impl->files.find(uri.file());
788   if (fileIt != impl->files.end())
789     return fileIt->second->findHover(uri, hoverPos);
790   return llvm::None;
791 }
792