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 ¬e : 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 ®istry, 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 ®istry, 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 ®ion) { 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 ®istry, 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 ®istry, 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 ®istry, 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 ®istry) : registry(registry) {} 735 736 /// The registry containing dialects that can be recognized in parsed .mlir 737 /// files. 738 DialectRegistry ®istry; 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 ®istry) 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