#![allow(clippy::mutable_key_type)]
use std::collections::HashMap;
use std::error::Error;
use lsp_server::{Connection, ExtractError, Message, Notification, Request, RequestId, Response};
use lsp_types::notification::{
DidChangeTextDocument, DidCloseTextDocument, DidOpenTextDocument, Notification as _,
PublishDiagnostics,
};
use lsp_types::request::{
FoldingRangeRequest, Formatting, HoverRequest, Request as _, SemanticTokensFullRequest,
};
use lsp_types::{
DidChangeTextDocumentParams, DidCloseTextDocumentParams, DidOpenTextDocumentParams,
DocumentFormattingParams, FoldingRange, FoldingRangeParams, Hover, HoverParams,
HoverProviderCapability, OneOf, PublishDiagnosticsParams, SemanticTokens, SemanticTokensLegend,
SemanticTokensOptions, SemanticTokensParams, SemanticTokensServerCapabilities,
ServerCapabilities, TextDocumentSyncCapability, TextDocumentSyncKind, Uri,
WorkDoneProgressOptions,
};
use sql_dialect_fmt_formatter::FormatOptions;
use sql_dialect_fmt_lsp::{
apply_change, diagnostics, folding_ranges, format_edits, hover, semantic_tokens,
token_modifiers, token_types,
};
type Docs = HashMap<Uri, String>;
fn main() -> Result<(), Box<dyn Error + Sync + Send>> {
let (connection, io_threads) = Connection::stdio();
let capabilities = serde_json::to_value(server_capabilities())?;
connection.initialize(capabilities)?;
main_loop(connection)?;
io_threads.join()?;
Ok(())
}
fn server_capabilities() -> ServerCapabilities {
ServerCapabilities {
text_document_sync: Some(TextDocumentSyncCapability::Kind(
TextDocumentSyncKind::INCREMENTAL,
)),
document_formatting_provider: Some(OneOf::Left(true)),
hover_provider: Some(HoverProviderCapability::Simple(true)),
folding_range_provider: Some(lsp_types::FoldingRangeProviderCapability::Simple(true)),
semantic_tokens_provider: Some(SemanticTokensServerCapabilities::SemanticTokensOptions(
SemanticTokensOptions {
legend: SemanticTokensLegend {
token_types: token_types(),
token_modifiers: token_modifiers(),
},
full: Some(lsp_types::SemanticTokensFullOptions::Bool(true)),
range: Some(false),
work_done_progress_options: WorkDoneProgressOptions::default(),
},
)),
..ServerCapabilities::default()
}
}
fn main_loop(connection: Connection) -> Result<(), Box<dyn Error + Sync + Send>> {
let mut docs = Docs::new();
for message in &connection.receiver {
match message {
Message::Request(request) => {
if connection.handle_shutdown(&request)? {
return Ok(());
}
let response = handle_request(request, &docs);
connection.sender.send(Message::Response(response))?;
}
Message::Notification(notification) => {
handle_notification(&connection, notification, &mut docs)?;
}
Message::Response(_) => {} }
}
Ok(())
}
fn handle_request(request: Request, docs: &Docs) -> Response {
match request.method.as_str() {
Formatting::METHOD => match cast::<Formatting>(request) {
Ok((id, params)) => ok(id, formatting(params, docs)),
Err(response) => *response,
},
SemanticTokensFullRequest::METHOD => match cast::<SemanticTokensFullRequest>(request) {
Ok((id, params)) => ok(id, semantic_tokens_full(params, docs)),
Err(response) => *response,
},
HoverRequest::METHOD => match cast::<HoverRequest>(request) {
Ok((id, params)) => ok(id, hover_request(params, docs)),
Err(response) => *response,
},
FoldingRangeRequest::METHOD => match cast::<FoldingRangeRequest>(request) {
Ok((id, params)) => ok(id, folding_request(params, docs)),
Err(response) => *response,
},
_ => Response::new_err(
request.id,
lsp_server::ErrorCode::MethodNotFound as i32,
format!("unsupported request: {}", request.method),
),
}
}
fn formatting(params: DocumentFormattingParams, docs: &Docs) -> Vec<lsp_types::TextEdit> {
let Some(text) = docs.get(¶ms.text_document.uri) else {
return Vec::new();
};
let options =
FormatOptions::default().with_indent_width((params.options.tab_size as usize).max(1));
format_edits(text, &options)
}
fn semantic_tokens_full(params: SemanticTokensParams, docs: &Docs) -> Option<SemanticTokens> {
let text = docs.get(¶ms.text_document.uri)?;
Some(SemanticTokens {
result_id: None,
data: semantic_tokens(text),
})
}
fn hover_request(params: HoverParams, docs: &Docs) -> Option<Hover> {
let position = params.text_document_position_params;
let text = docs.get(&position.text_document.uri)?;
hover(text, position.position)
}
fn folding_request(params: FoldingRangeParams, docs: &Docs) -> Option<Vec<FoldingRange>> {
let text = docs.get(¶ms.text_document.uri)?;
Some(folding_ranges(text))
}
fn handle_notification(
connection: &Connection,
notification: Notification,
docs: &mut Docs,
) -> Result<(), Box<dyn Error + Sync + Send>> {
match notification.method.as_str() {
DidOpenTextDocument::METHOD => {
let params: DidOpenTextDocumentParams =
notification.extract(DidOpenTextDocument::METHOD)?;
let uri = params.text_document.uri;
docs.insert(uri.clone(), params.text_document.text);
publish_diagnostics(connection, docs, &uri)?;
}
DidChangeTextDocument::METHOD => {
let params: DidChangeTextDocumentParams =
notification.extract(DidChangeTextDocument::METHOD)?;
let uri = params.text_document.uri;
let mut text = docs.remove(&uri).unwrap_or_default();
for change in params.content_changes {
text = apply_change(&text, change.range, &change.text);
}
docs.insert(uri.clone(), text);
publish_diagnostics(connection, docs, &uri)?;
}
DidCloseTextDocument::METHOD => {
let params: DidCloseTextDocumentParams =
notification.extract(DidCloseTextDocument::METHOD)?;
let uri = params.text_document.uri;
docs.remove(&uri);
send_diagnostics(connection, uri, Vec::new())?; }
_ => {}
}
Ok(())
}
fn publish_diagnostics(
connection: &Connection,
docs: &Docs,
uri: &Uri,
) -> Result<(), Box<dyn Error + Sync + Send>> {
let diags = docs
.get(uri)
.map(|text| diagnostics(text))
.unwrap_or_default();
send_diagnostics(connection, uri.clone(), diags)
}
fn send_diagnostics(
connection: &Connection,
uri: Uri,
diagnostics: Vec<lsp_types::Diagnostic>,
) -> Result<(), Box<dyn Error + Sync + Send>> {
let params = PublishDiagnosticsParams {
uri,
diagnostics,
version: None,
};
let notification = Notification::new(PublishDiagnostics::METHOD.to_string(), params);
connection
.sender
.send(Message::Notification(notification))?;
Ok(())
}
fn ok<T: serde::Serialize>(id: RequestId, result: T) -> Response {
Response::new_ok(id, result)
}
fn cast<R>(request: Request) -> Result<(RequestId, R::Params), Box<Response>>
where
R: lsp_types::request::Request,
{
request.extract(R::METHOD).map_err(|err| match err {
ExtractError::JsonError { method, error } => Box::new(Response::new_err(
RequestId::from(0),
lsp_server::ErrorCode::InvalidParams as i32,
format!("invalid params for {method}: {error}"),
)),
ExtractError::MethodMismatch(request) => Box::new(Response::new_err(
request.id,
lsp_server::ErrorCode::MethodNotFound as i32,
format!("method mismatch: {}", request.method),
)),
})
}