381 lines · cpp
1//===- LSPServer.cpp - MLIR 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-exception6//7//===----------------------------------------------------------------------===//8 9#include "LSPServer.h"10#include "MLIRServer.h"11#include "Protocol.h"12#include "llvm/Support/LSP/Logging.h"13#include "llvm/Support/LSP/Transport.h"14#include <optional>15 16#define DEBUG_TYPE "mlir-lsp-server"17 18using namespace mlir;19using namespace mlir::lsp;20 21using llvm::lsp::Callback;22using llvm::lsp::CodeAction;23using llvm::lsp::CodeActionParams;24using llvm::lsp::CompletionList;25using llvm::lsp::CompletionParams;26using llvm::lsp::DidChangeTextDocumentParams;27using llvm::lsp::DidCloseTextDocumentParams;28using llvm::lsp::DidOpenTextDocumentParams;29using llvm::lsp::DocumentSymbol;30using llvm::lsp::DocumentSymbolParams;31using llvm::lsp::Hover;32using llvm::lsp::InitializedParams;33using llvm::lsp::InitializeParams;34using llvm::lsp::JSONTransport;35using llvm::lsp::Location;36using llvm::lsp::Logger;37using llvm::lsp::MessageHandler;38using llvm::lsp::MLIRConvertBytecodeParams;39using llvm::lsp::MLIRConvertBytecodeResult;40using llvm::lsp::NoParams;41using llvm::lsp::OutgoingNotification;42using llvm::lsp::PublishDiagnosticsParams;43using llvm::lsp::ReferenceParams;44using llvm::lsp::TextDocumentPositionParams;45using llvm::lsp::TextDocumentSyncKind;46using llvm::lsp::URIForFile;47 48//===----------------------------------------------------------------------===//49// LSPServer50//===----------------------------------------------------------------------===//51 52namespace {53struct LSPServer {54 LSPServer(MLIRServer &server) : server(server) {}55 56 //===--------------------------------------------------------------------===//57 // Initialization58 59 void onInitialize(const InitializeParams ¶ms,60 Callback<llvm::json::Value> reply);61 void onInitialized(const InitializedParams ¶ms);62 void onShutdown(const NoParams ¶ms, Callback<std::nullptr_t> reply);63 64 //===--------------------------------------------------------------------===//65 // Document Change66 67 void onDocumentDidOpen(const DidOpenTextDocumentParams ¶ms);68 void onDocumentDidClose(const DidCloseTextDocumentParams ¶ms);69 void onDocumentDidChange(const DidChangeTextDocumentParams ¶ms);70 71 //===--------------------------------------------------------------------===//72 // Definitions and References73 74 void onGoToDefinition(const TextDocumentPositionParams ¶ms,75 Callback<std::vector<Location>> reply);76 void onReference(const ReferenceParams ¶ms,77 Callback<std::vector<Location>> reply);78 79 //===--------------------------------------------------------------------===//80 // Hover81 82 void onHover(const TextDocumentPositionParams ¶ms,83 Callback<std::optional<Hover>> reply);84 85 //===--------------------------------------------------------------------===//86 // Document Symbols87 88 void onDocumentSymbol(const DocumentSymbolParams ¶ms,89 Callback<std::vector<DocumentSymbol>> reply);90 91 //===--------------------------------------------------------------------===//92 // Code Completion93 94 void onCompletion(const CompletionParams ¶ms,95 Callback<CompletionList> reply);96 97 //===--------------------------------------------------------------------===//98 // Code Action99 100 void onCodeAction(const CodeActionParams ¶ms,101 Callback<llvm::json::Value> reply);102 103 //===--------------------------------------------------------------------===//104 // Bytecode105 106 void onConvertFromBytecode(const MLIRConvertBytecodeParams ¶ms,107 Callback<MLIRConvertBytecodeResult> reply);108 void onConvertToBytecode(const MLIRConvertBytecodeParams ¶ms,109 Callback<MLIRConvertBytecodeResult> reply);110 111 //===--------------------------------------------------------------------===//112 // Fields113 //===--------------------------------------------------------------------===//114 115 MLIRServer &server;116 117 /// An outgoing notification used to send diagnostics to the client when they118 /// are ready to be processed.119 OutgoingNotification<PublishDiagnosticsParams> publishDiagnostics;120 121 /// Used to indicate that the 'shutdown' request was received from the122 /// Language Server client.123 bool shutdownRequestReceived = false;124};125} // namespace126 127//===----------------------------------------------------------------------===//128// Initialization129//===----------------------------------------------------------------------===//130 131void LSPServer::onInitialize(const InitializeParams ¶ms,132 Callback<llvm::json::Value> reply) {133 // Send a response with the capabilities of this server.134 llvm::json::Object serverCaps{135 {"textDocumentSync",136 llvm::json::Object{137 {"openClose", true},138 {"change", (int)TextDocumentSyncKind::Full},139 {"save", true},140 }},141 {"completionProvider",142 llvm::json::Object{143 {"allCommitCharacters",144 {145 "\t",146 ";",147 ",",148 ".",149 "=",150 }},151 {"resolveProvider", false},152 {"triggerCharacters",153 {".", "%", "^", "!", "#", "(", ",", "<", ":", "[", " ", "\"", "/"}},154 }},155 {"definitionProvider", true},156 {"referencesProvider", true},157 {"hoverProvider", true},158 159 // For now we only support documenting symbols when the client supports160 // hierarchical symbols.161 {"documentSymbolProvider",162 params.capabilities.hierarchicalDocumentSymbol},163 };164 165 // Per LSP, codeActionProvider can be either boolean or CodeActionOptions.166 // CodeActionOptions is only valid if the client supports action literal167 // via textDocument.codeAction.codeActionLiteralSupport.168 serverCaps["codeActionProvider"] =169 params.capabilities.codeActionStructure170 ? llvm::json::Object{{"codeActionKinds",171 {CodeAction::kQuickFix, CodeAction::kRefactor,172 CodeAction::kInfo}}}173 : llvm::json::Value(true);174 175 llvm::json::Object result{176 {{"serverInfo",177 llvm::json::Object{{"name", "mlir-lsp-server"}, {"version", "0.0.0"}}},178 {"capabilities", std::move(serverCaps)}}};179 reply(std::move(result));180}181void LSPServer::onInitialized(const InitializedParams &) {}182void LSPServer::onShutdown(const NoParams &, Callback<std::nullptr_t> reply) {183 shutdownRequestReceived = true;184 reply(nullptr);185}186 187//===----------------------------------------------------------------------===//188// Document Change189//===----------------------------------------------------------------------===//190 191void LSPServer::onDocumentDidOpen(const DidOpenTextDocumentParams ¶ms) {192 PublishDiagnosticsParams diagParams(params.textDocument.uri,193 params.textDocument.version);194 server.addOrUpdateDocument(params.textDocument.uri, params.textDocument.text,195 params.textDocument.version,196 diagParams.diagnostics);197 198 // Publish any recorded diagnostics.199 publishDiagnostics(diagParams);200}201void LSPServer::onDocumentDidClose(const DidCloseTextDocumentParams ¶ms) {202 std::optional<int64_t> version =203 server.removeDocument(params.textDocument.uri);204 if (!version)205 return;206 207 // Empty out the diagnostics shown for this document. This will clear out208 // anything currently displayed by the client for this document (e.g. in the209 // "Problems" pane of VSCode).210 publishDiagnostics(211 PublishDiagnosticsParams(params.textDocument.uri, *version));212}213void LSPServer::onDocumentDidChange(const DidChangeTextDocumentParams ¶ms) {214 // TODO: We currently only support full document updates, we should refactor215 // to avoid this.216 if (params.contentChanges.size() != 1)217 return;218 PublishDiagnosticsParams diagParams(params.textDocument.uri,219 params.textDocument.version);220 server.addOrUpdateDocument(221 params.textDocument.uri, params.contentChanges.front().text,222 params.textDocument.version, diagParams.diagnostics);223 224 // Publish any recorded diagnostics.225 publishDiagnostics(diagParams);226}227 228//===----------------------------------------------------------------------===//229// Definitions and References230//===----------------------------------------------------------------------===//231 232void LSPServer::onGoToDefinition(const TextDocumentPositionParams ¶ms,233 Callback<std::vector<Location>> reply) {234 std::vector<Location> locations;235 server.getLocationsOf(params.textDocument.uri, params.position, locations);236 reply(std::move(locations));237}238 239void LSPServer::onReference(const ReferenceParams ¶ms,240 Callback<std::vector<Location>> reply) {241 std::vector<Location> locations;242 server.findReferencesOf(params.textDocument.uri, params.position, locations);243 reply(std::move(locations));244}245 246//===----------------------------------------------------------------------===//247// Hover248//===----------------------------------------------------------------------===//249 250void LSPServer::onHover(const TextDocumentPositionParams ¶ms,251 Callback<std::optional<Hover>> reply) {252 reply(server.findHover(params.textDocument.uri, params.position));253}254 255//===----------------------------------------------------------------------===//256// Document Symbols257//===----------------------------------------------------------------------===//258 259void LSPServer::onDocumentSymbol(const DocumentSymbolParams ¶ms,260 Callback<std::vector<DocumentSymbol>> reply) {261 std::vector<DocumentSymbol> symbols;262 server.findDocumentSymbols(params.textDocument.uri, symbols);263 reply(std::move(symbols));264}265 266//===----------------------------------------------------------------------===//267// Code Completion268//===----------------------------------------------------------------------===//269 270void LSPServer::onCompletion(const CompletionParams ¶ms,271 Callback<CompletionList> reply) {272 reply(server.getCodeCompletion(params.textDocument.uri, params.position));273}274 275//===----------------------------------------------------------------------===//276// Code Action277//===----------------------------------------------------------------------===//278 279void LSPServer::onCodeAction(const CodeActionParams ¶ms,280 Callback<llvm::json::Value> reply) {281 URIForFile uri = params.textDocument.uri;282 283 // Check whether a particular CodeActionKind is included in the response.284 auto isKindAllowed = [only(params.context.only)](StringRef kind) {285 if (only.empty())286 return true;287 return llvm::any_of(only, [&](StringRef base) {288 return kind.consume_front(base) &&289 (kind.empty() || kind.starts_with("."));290 });291 };292 293 // We provide a code action for fixes on the specified diagnostics.294 std::vector<CodeAction> actions;295 if (isKindAllowed(CodeAction::kQuickFix))296 server.getCodeActions(uri, params.range.start, params.context, actions);297 reply(std::move(actions));298}299 300//===----------------------------------------------------------------------===//301// Bytecode302//===----------------------------------------------------------------------===//303 304void LSPServer::onConvertFromBytecode(305 const MLIRConvertBytecodeParams ¶ms,306 Callback<MLIRConvertBytecodeResult> reply) {307 reply(server.convertFromBytecode(params.uri));308}309 310void LSPServer::onConvertToBytecode(const MLIRConvertBytecodeParams ¶ms,311 Callback<MLIRConvertBytecodeResult> reply) {312 reply(server.convertToBytecode(params.uri));313}314 315//===----------------------------------------------------------------------===//316// Entry point317//===----------------------------------------------------------------------===//318 319LogicalResult lsp::runMlirLSPServer(MLIRServer &server,320 JSONTransport &transport) {321 LSPServer lspServer(server);322 MessageHandler messageHandler(transport);323 324 // Initialization325 messageHandler.method("initialize", &lspServer, &LSPServer::onInitialize);326 messageHandler.notification("initialized", &lspServer,327 &LSPServer::onInitialized);328 messageHandler.method("shutdown", &lspServer, &LSPServer::onShutdown);329 330 // Document Changes331 messageHandler.notification("textDocument/didOpen", &lspServer,332 &LSPServer::onDocumentDidOpen);333 messageHandler.notification("textDocument/didClose", &lspServer,334 &LSPServer::onDocumentDidClose);335 messageHandler.notification("textDocument/didChange", &lspServer,336 &LSPServer::onDocumentDidChange);337 338 // Definitions and References339 messageHandler.method("textDocument/definition", &lspServer,340 &LSPServer::onGoToDefinition);341 messageHandler.method("textDocument/references", &lspServer,342 &LSPServer::onReference);343 344 // Hover345 messageHandler.method("textDocument/hover", &lspServer, &LSPServer::onHover);346 347 // Document Symbols348 messageHandler.method("textDocument/documentSymbol", &lspServer,349 &LSPServer::onDocumentSymbol);350 351 // Code Completion352 messageHandler.method("textDocument/completion", &lspServer,353 &LSPServer::onCompletion);354 355 // Code Action356 messageHandler.method("textDocument/codeAction", &lspServer,357 &LSPServer::onCodeAction);358 359 // Bytecode360 messageHandler.method("mlir/convertFromBytecode", &lspServer,361 &LSPServer::onConvertFromBytecode);362 messageHandler.method("mlir/convertToBytecode", &lspServer,363 &LSPServer::onConvertToBytecode);364 365 // Diagnostics366 lspServer.publishDiagnostics =367 messageHandler.outgoingNotification<PublishDiagnosticsParams>(368 "textDocument/publishDiagnostics");369 370 // Run the main loop of the transport.371 LogicalResult result = success();372 if (llvm::Error error = transport.run(messageHandler)) {373 Logger::error("Transport error: {0}", error);374 llvm::consumeError(std::move(error));375 result = failure();376 } else {377 result = success(lspServer.shutdownRequestReceived);378 }379 return result;380}381