package lsp import ( "bytes" "context" "encoding/json" "fmt" "strings" "testing" "time" ) // rpc sends a JSON-RPC request frame and returns the raw body bytes. func frame(id int, method string, params interface{}) []byte { type req struct { JSONRPC string `json:"jsonrpc"` ID int `json:"id"` Method string `json:"method"` Params interface{} `json:"params"` } body, _ := json.Marshal(req{JSONRPC: "2.0", ID: id, Method: method, Params: params}) return []byte(fmt.Sprintf("Content-Length: %d\r\n\r\n%s", len(body), body)) } func notifFrame(method string, params interface{}) []byte { type notif struct { JSONRPC string `json:"jsonrpc"` Method string `json:"method"` Params interface{} `json:"params"` } body, _ := json.Marshal(notif{JSONRPC: "2.0", Method: method, Params: params}) return []byte(fmt.Sprintf("Content-Length: %d\r\n\r\n%s", len(body), body)) } // readResp reads one JSON-RPC response from a *bytes.Buffer (blocking until available). func readResp(t *testing.T, buf *bytes.Buffer) map[string]interface{} { t.Helper() deadline := time.Now().Add(2 * time.Second) for time.Now().Before(deadline) { data := buf.Bytes() // Find Content-Length header idx := bytes.Index(data, []byte("Content-Length: ")) if idx < 0 { time.Sleep(5 * time.Millisecond) continue } eol := bytes.Index(data[idx:], []byte("\r\n")) if eol < 0 { time.Sleep(5 * time.Millisecond) continue } lenStr := string(data[idx+16 : idx+eol]) var n int _, _ = fmt.Sscanf(lenStr, "%d", &n) sep := bytes.Index(data, []byte("\r\n\r\n")) if sep < 0 || len(data) < sep+4+n { time.Sleep(5 * time.Millisecond) continue } body := data[sep+4 : sep+4+n] buf.Next(sep + 4 + n) var result map[string]interface{} _ = json.Unmarshal(body, &result) return result } t.Fatal("timeout waiting for response") return nil } func runServer(t *testing.T, input []byte) *bytes.Buffer { t.Helper() out := &bytes.Buffer{} ctx, cancel := context.WithTimeout(context.Background(), 3*time.Second) done := make(chan error, 1) go func() { done <- Serve(ctx, bytes.NewReader(input), out, t.TempDir()) cancel() }() // Give the server time to process select { case <-done: case <-time.After(3 * time.Second): } return out } func TestInitialize(t *testing.T) { input := append(frame(1, "initialize", map[string]interface{}{}), notifFrame("initialized", map[string]interface{}{})...) input = append(input, frame(2, "shutdown", nil)...) input = append(input, notifFrame("exit", nil)...) out := runServer(t, input) resp := readResp(t, out) if resp["id"].(float64) != 1 { t.Fatalf("expected id=1, got %v", resp["id"]) } caps := resp["result"].(map[string]interface{})["capabilities"].(map[string]interface{}) if caps["documentFormattingProvider"] != true { t.Error("expected documentFormattingProvider=true") } } func TestFormatting(t *testing.T) { unformatted := "create function foo() returns void language plpgsql as $$ begin end $$;" uri := "file:///test.sql" var input []byte input = append(input, frame(1, "initialize", map[string]interface{}{})...) input = append(input, notifFrame("initialized", nil)...) input = append(input, notifFrame("textDocument/didOpen", map[string]interface{}{ "textDocument": map[string]interface{}{ "uri": uri, "languageId": "sql", "version": 1, "text": unformatted, }, })...) input = append(input, frame(2, "textDocument/formatting", map[string]interface{}{ "textDocument": map[string]interface{}{"uri": uri}, "options": map[string]interface{}{"tabSize": 2, "insertSpaces": true}, })...) input = append(input, frame(3, "shutdown", nil)...) input = append(input, notifFrame("exit", nil)...) out := runServer(t, input) // Skip initialize response and publishDiagnostics notification, find formatting response var formattingResp map[string]interface{} for i := 0; i < 10; i++ { resp := readResp(t, out) if resp == nil { break } id, hasID := resp["id"] if hasID && id.(float64) == 2 { formattingResp = resp break } } if formattingResp == nil { t.Fatal("did not receive formatting response") } result, ok := formattingResp["result"].([]interface{}) if !ok { t.Fatalf("expected array result, got %T: %v", formattingResp["result"], formattingResp["result"]) } if len(result) == 0 { t.Fatal("expected at least one text edit") } edit := result[0].(map[string]interface{}) newText := edit["newText"].(string) if !strings.Contains(newText, "CREATE FUNCTION") { t.Errorf("formatted output missing CREATE FUNCTION keyword; got:\n%s", newText) } } func TestRangeFormatting(t *testing.T) { // Two statements: first is unformatted, second is already formatted. // Range request covers only the first statement (line 0). sql := "select a, b from t where x = 1;\nSELECT z\nFROM s;\n" uri := "file:///range.sql" var input []byte input = append(input, frame(1, "initialize", map[string]interface{}{})...) input = append(input, notifFrame("initialized", nil)...) input = append(input, notifFrame("textDocument/didOpen", map[string]interface{}{ "textDocument": map[string]interface{}{ "uri": uri, "languageId": "sql", "version": 1, "text": sql, }, })...) input = append(input, frame(2, "textDocument/rangeFormatting", map[string]interface{}{ "textDocument": map[string]interface{}{"uri": uri}, "range": map[string]interface{}{ "start": map[string]interface{}{"line": 0, "character": 0}, "end": map[string]interface{}{"line": 0, "character": 31}, }, "options": map[string]interface{}{"tabSize": 2, "insertSpaces": true}, })...) input = append(input, frame(3, "shutdown", nil)...) input = append(input, notifFrame("exit", nil)...) out := runServer(t, input) var resp map[string]interface{} for i := 0; i < 10; i++ { r := readResp(t, out) if r == nil { break } if id, ok := r["id"]; ok && id.(float64) == 2 { resp = r break } } if resp == nil { t.Fatal("no response for rangeFormatting request") } result, ok := resp["result"].([]interface{}) if !ok || len(result) == 0 { t.Fatalf("expected non-empty edit array, got %v", resp["result"]) } edit := result[0].(map[string]interface{}) newText, _ := edit["newText"].(string) if !strings.Contains(newText, "SELECT") { t.Errorf("expected formatted SELECT in edit, got: %q", newText) } if strings.Contains(newText, "SELECT z") { t.Errorf("range formatting edited second statement unexpectedly; got: %q", newText) } } func TestInitialize_AdvertisesRangeFormatting(t *testing.T) { input := append(frame(1, "initialize", map[string]interface{}{}), notifFrame("initialized", map[string]interface{}{})...) input = append(input, frame(2, "shutdown", nil)...) input = append(input, notifFrame("exit", nil)...) out := runServer(t, input) resp := readResp(t, out) caps := resp["result"].(map[string]interface{})["capabilities"].(map[string]interface{}) if caps["documentRangeFormattingProvider"] != true { t.Error("expected documentRangeFormattingProvider=true") } } func TestDidClose_ClearsDiagnostics(t *testing.T) { uri := "file:///test.sql" // SQL with a lint violation (SELECT * triggers COR001) sql := "SELECT * FROM users;" var input []byte input = append(input, frame(1, "initialize", map[string]interface{}{})...) input = append(input, notifFrame("initialized", nil)...) input = append(input, notifFrame("textDocument/didOpen", map[string]interface{}{ "textDocument": map[string]interface{}{ "uri": uri, "languageId": "sql", "version": 1, "text": sql, }, })...) input = append(input, notifFrame("textDocument/didClose", map[string]interface{}{ "textDocument": map[string]interface{}{"uri": uri}, })...) input = append(input, frame(2, "shutdown", nil)...) input = append(input, notifFrame("exit", nil)...) out := runServer(t, input) // Collect all publishDiagnostics notifications; the last one for this URI must be empty. var lastDiags []interface{} for i := 0; i < 20; i++ { if out.Len() == 0 { break } resp := readResp(t, out) if resp == nil { break } if resp["method"] == "textDocument/publishDiagnostics" { p := resp["params"].(map[string]interface{}) if p["uri"] == uri { lastDiags = p["diagnostics"].([]interface{}) } } } if len(lastDiags) != 0 { t.Errorf("expected empty diagnostics after didClose, got %d", len(lastDiags)) } } // request opens text at uri, sends one request, and returns its response // (skipping the initialize response and any notifications). func request(t *testing.T, text, method string, params map[string]interface{}) map[string]interface{} { t.Helper() uri := "file:///t.sql" params["textDocument"] = map[string]interface{}{"uri": uri} var input []byte input = append(input, frame(1, "initialize", map[string]interface{}{})...) input = append(input, notifFrame("initialized", nil)...) input = append(input, notifFrame("textDocument/didOpen", map[string]interface{}{ "textDocument": map[string]interface{}{"uri": uri, "languageId": "sql", "version": 1, "text": text}, })...) input = append(input, frame(2, method, params)...) input = append(input, frame(3, "shutdown", nil)...) input = append(input, notifFrame("exit", nil)...) out := runServer(t, input) for i := 0; i < 10; i++ { resp := readResp(t, out) if id, ok := resp["id"]; ok && id.(float64) == 2 { return resp } } t.Fatalf("no response to %s", method) return nil } func TestDocumentSymbol(t *testing.T) { src := "create function public.foo(a int) returns void language plpgsql as $$ begin end $$;\n\ncreate procedure bar() language plpgsql as $$ begin end $$;\n" resp := request(t, src, "textDocument/documentSymbol", map[string]interface{}{}) syms := resp["result"].([]interface{}) if len(syms) != 2 { t.Fatalf("want 2 symbols, got %v", syms) } first := syms[0].(map[string]interface{}) if first["name"] != "public.foo" || first["detail"] != "function" { t.Errorf("first symbol: %v", first) } second := syms[1].(map[string]interface{}) if second["name"] != "bar" || second["detail"] != "procedure" { t.Errorf("second symbol: %v", second) } line := second["range"].(map[string]interface{})["start"].(map[string]interface{})["line"].(float64) if line != 2 { t.Errorf("bar should start on line 2, got %v", line) } } func TestHoverShowsDiagnostic(t *testing.T) { resp := request(t, "select * from t;", "textDocument/hover", map[string]interface{}{ "position": map[string]interface{}{"line": 0, "character": 7}, }) res, ok := resp["result"].(map[string]interface{}) if !ok { t.Fatalf("expected hover result, got %v", resp["result"]) } val := res["contents"].(map[string]interface{})["value"].(string) if !strings.Contains(val, "COR001") { t.Errorf("hover should name the rule, got %q", val) } resp = request(t, "select * from t;", "textDocument/hover", map[string]interface{}{ "position": map[string]interface{}{"line": 0, "character": 14}, }) if resp["result"] != nil { t.Errorf("no diagnostic at col 14, want null, got %v", resp["result"]) } } func TestWillSaveWaitUntilFormats(t *testing.T) { resp := request(t, "select a from t", "textDocument/willSaveWaitUntil", map[string]interface{}{"reason": 1}) edits := resp["result"].([]interface{}) if len(edits) != 1 || !strings.Contains(edits[0].(map[string]interface{})["newText"].(string), "SELECT a") { t.Errorf("expected a formatting edit, got %v", edits) } } func TestDiagSpanCoversToken(t *testing.T) { if got := diagSpan("select * from t;", 0, 7); got != 1 { t.Errorf("'*' span = %d, want 1", got) } if got := diagSpan("select foo from t;", 0, 7); got != 3 { t.Errorf("'foo' span = %d, want 3", got) } if got := diagSpan("select foo from t;", 0, 40); got != 1 { t.Errorf("out-of-range span = %d, want 1", got) } }