diff --git a/internal/ls/selectionranges.go b/internal/ls/selectionranges.go index 834f250fbb0..ec8fedbf453 100644 --- a/internal/ls/selectionranges.go +++ b/internal/ls/selectionranges.go @@ -10,6 +10,41 @@ import ( "github.com/microsoft/typescript-go/internal/scanner" ) +const maxSelectionRangeDepth = 1000 + +type selectionRangeBuilder struct { + ranges []lsproto.Range + oldestIndex int +} + +func newSelectionRangeBuilder(capacity int) *selectionRangeBuilder { + return &selectionRangeBuilder{ + ranges: make([]lsproto.Range, 0, capacity), + } +} + +func (b *selectionRangeBuilder) push(selectionRange lsproto.Range) { + if len(b.ranges) < cap(b.ranges) { + b.ranges = append(b.ranges, selectionRange) + return + } + + b.ranges[b.oldestIndex] = selectionRange + b.oldestIndex = (b.oldestIndex + 1) % len(b.ranges) +} + +func (b *selectionRangeBuilder) build(parentRange lsproto.Range) *lsproto.SelectionRange { + result := &lsproto.SelectionRange{Range: parentRange} + for i := range b.ranges { + index := (b.oldestIndex + i) % len(b.ranges) + result = &lsproto.SelectionRange{ + Range: b.ranges[index], + Parent: result, + } + } + return result +} + func (l *LanguageService) ProvideSelectionRanges(ctx context.Context, params *lsproto.SelectionRangeParams) (lsproto.SelectionRangeResponse, error) { _, sourceFile := l.getProgramAndFile(params.TextDocument.Uri) if sourceFile == nil { @@ -146,6 +181,10 @@ func createSyntaxList(factory *ast.NodeFactory, children []*ast.Node) *ast.Node func getSmartSelectionRange(l *LanguageService, sourceFile *ast.SourceFile, pos int) *lsproto.SelectionRange { factory := &ast.NodeFactory{} + fullRange := l.converters.ToLSPRange(sourceFile, core.NewTextRange(sourceFile.Pos(), sourceFile.End())) + // Traversal discovers ranges from broadest to most specific, so retain the newest ranges nearest to the cursor + ranges := newSelectionRangeBuilder(maxSelectionRangeDepth - 1) + lastRange := fullRange nodeContainsPosition := func(node *ast.Node) bool { if node == nil { @@ -167,38 +206,34 @@ func getSmartSelectionRange(l *LanguageService, sourceFile *ast.SourceFile, pos return false } - pushSelectionRange := func(current *lsproto.SelectionRange, start, end int) *lsproto.SelectionRange { + pushSelectionRange := func(start, end int) { if start == end { - return current + return } if !(start <= pos && pos <= end) { - return current + return } lspRange := l.converters.ToLSPRange(sourceFile, core.NewTextRange(start, end)) - if current != nil && current.Range == lspRange { - return current + if lastRange == lspRange { + return } + lastRange = lspRange - return &lsproto.SelectionRange{ - Range: lspRange, - Parent: current, - } + ranges.push(lspRange) } - pushSelectionCommentRange := func(current *lsproto.SelectionRange, start, end int) *lsproto.SelectionRange { - current = pushSelectionRange(current, start, end) + pushSelectionCommentRange := func(start, end int) { + pushSelectionRange(start, end) commentPos := start text := sourceFile.Text() for commentPos < end && commentPos < len(text) && text[commentPos] == '/' { commentPos++ } - current = pushSelectionRange(current, commentPos, end) - - return current + pushSelectionRange(commentPos, end) } positionsAreOnSameLine := func(pos1, pos2 int) bool { @@ -238,11 +273,6 @@ func getSmartSelectionRange(l *LanguageService, sourceFile *ast.SourceFile, pos return false } - fullRange := l.converters.ToLSPRange(sourceFile, core.NewTextRange(sourceFile.Pos(), sourceFile.End())) - result := &lsproto.SelectionRange{ - Range: fullRange, - } - var current *ast.Node for current = sourceFile.AsNode(); current != nil; { var next *ast.Node @@ -256,7 +286,7 @@ func getSmartSelectionRange(l *LanguageService, sourceFile *ast.SourceFile, pos break } if foundComment != nil && foundComment.Kind == ast.KindSingleLineCommentTrivia { - result = pushSelectionCommentRange(result, foundComment.Pos(), foundComment.End()) + pushSelectionCommentRange(foundComment.Pos(), foundComment.End()) } if nodeContainsPosition(node) { @@ -265,7 +295,7 @@ func getSmartSelectionRange(l *LanguageService, sourceFile *ast.SourceFile, pos if !positionsAreOnSameLine(astnav.GetStartOfNode(node, sourceFile, false), node.End()) { start := astnav.GetStartOfNode(node, sourceFile, false) end := node.End() - result = pushSelectionRange(result, start, end) + pushSelectionRange(start, end) } } @@ -281,7 +311,7 @@ func getSmartSelectionRange(l *LanguageService, sourceFile *ast.SourceFile, pos // Validate the positions are reasonable text := sourceFile.Text() if spanStart >= 0 && spanEnd <= len(text) && spanStart < spanEnd { - result = pushSelectionRange(result, spanStart, spanEnd) + pushSelectionRange(spanStart, spanEnd) } } } @@ -289,7 +319,7 @@ func getSmartSelectionRange(l *LanguageService, sourceFile *ast.SourceFile, pos if !shouldSkipNode(node, parent) { start := astnav.GetStartOfNode(node, sourceFile, false) end := node.End() - result = pushSelectionRange(result, start, end) + pushSelectionRange(start, end) if ast.IsMappedTypeNode(node) { for selectionParent := node; ; { @@ -300,7 +330,7 @@ func getSmartSelectionRange(l *LanguageService, sourceFile *ast.SourceFile, pos break } if positionShouldSnapToNode(child) { - result = pushSelectionRange(result, childStart, child.End()) + pushSelectionRange(childStart, child.End()) selectionChild = child break } @@ -316,7 +346,7 @@ func getSmartSelectionRange(l *LanguageService, sourceFile *ast.SourceFile, pos if ast.IsStringLiteral(node) || node.Kind == ast.KindTemplateExpression || node.Kind == ast.KindNoSubstitutionTemplateLiteral { // Only add inner content range if there's actually content (handles unterminated literals) if start+1 < end-1 { - result = pushSelectionRange(result, start+1, end-1) + pushSelectionRange(start+1, end-1) } } } @@ -336,7 +366,7 @@ func getSmartSelectionRange(l *LanguageService, sourceFile *ast.SourceFile, pos end := nodes.Nodes[len(nodes.Nodes)-1].End() if start <= pos && pos < end { - result = pushSelectionRange(result, start, end) + pushSelectionRange(start, end) } } } @@ -355,5 +385,5 @@ func getSmartSelectionRange(l *LanguageService, sourceFile *ast.SourceFile, pos current.VisitEachChild(tempVisitor) current = next } - return result + return ranges.build(fullRange) } diff --git a/internal/ls/selectionranges_test.go b/internal/ls/selectionranges_test.go new file mode 100644 index 00000000000..48a41ec425d --- /dev/null +++ b/internal/ls/selectionranges_test.go @@ -0,0 +1,58 @@ +package ls + +import ( + "strings" + "testing" + + "github.com/microsoft/typescript-go/internal/ast" + "github.com/microsoft/typescript-go/internal/core" + "github.com/microsoft/typescript-go/internal/json" + "github.com/microsoft/typescript-go/internal/jsonrpc" + "github.com/microsoft/typescript-go/internal/ls/lsconv" + "github.com/microsoft/typescript-go/internal/lsp/lsproto" + "github.com/microsoft/typescript-go/internal/parser" +) + +func TestSelectionRangeDepthIsLimited(t *testing.T) { + t.Parallel() + + const nestingDepth = 12000 + text := "const x = " + strings.Repeat("(", nestingDepth) + "1" + strings.Repeat(")", nestingDepth) + ";" + sourceFile := parser.ParseSourceFile(ast.SourceFileParseOptions{ + FileName: "/index.ts", + Path: "/index.ts", + }, text, core.ScriptKindTS) + lineMap := lsconv.ComputeLSPLineStarts(text) + languageService := &LanguageService{ + converters: lsconv.NewConverters(lsproto.PositionEncodingKindUTF16, func(string) *lsconv.LSPLineMap { + return lineMap + }), + } + + result := getSmartSelectionRange(languageService, sourceFile, len("const x = ")+nestingDepth) + depth := 0 + var outermost *lsproto.SelectionRange + for current := result; current != nil; current = current.Parent { + depth++ + outermost = current + } + + if depth != maxSelectionRangeDepth { + t.Fatalf("selection range depth = %d, want %d", depth, maxSelectionRangeDepth) + } + innerRange := languageService.converters.ToLSPRange(sourceFile, core.NewTextRange(len("const x = ")+nestingDepth, len("const x = ")+nestingDepth+1)) + if result.Range != innerRange { + t.Fatalf("innermost selection range = %v, want %v", result.Range, innerRange) + } + fullRange := languageService.converters.ToLSPRange(sourceFile, core.NewTextRange(sourceFile.Pos(), sourceFile.End())) + if outermost.Range != fullRange { + t.Fatalf("outermost selection range = %v, want full file range %v", outermost.Range, fullRange) + } + results := []*lsproto.SelectionRange{result} + response := lsproto.SelectionRangesOrNull{SelectionRanges: &results} + id := jsonrpc.NewIDString("selectionRange") + message := (&lsproto.ResponseMessage{ID: id, Result: &response}).Message() + if _, err := json.Marshal(message); err != nil { + t.Fatalf("failed to marshal limited selection range: %v", err) + } +} diff --git a/internal/lsp/server.go b/internal/lsp/server.go index 5ba37e85880..0ded229e354 100644 --- a/internal/lsp/server.go +++ b/internal/lsp/server.go @@ -113,6 +113,16 @@ type lspWriter struct { w *lsproto.BaseWriter } +type messageMarshalError struct { + err error +} + +func (e *messageMarshalError) Error() string { return "failed to marshal message: " + e.err.Error() } + +func (e *messageMarshalError) Unwrap() []error { + return []error{lsproto.ErrorCodeInternalError, e.err} +} + func (r *lspReader) Read() (*lsproto.Message, error) { data, err := r.r.Read() if err != nil { @@ -137,7 +147,7 @@ func ToReader(r io.Reader) Reader { func (w *lspWriter) Write(msg *lsproto.Message) error { data, err := json.Marshal(msg) if err != nil { - return fmt.Errorf("failed to marshal message: %w", err) + return &messageMarshalError{err: err} } return w.w.Write(data) } @@ -626,6 +636,16 @@ func (s *Server) writeLoop(ctx context.Context) error { return err } if err := s.w.Write(msg); err != nil { + var marshalErr *messageMarshalError + if errors.As(err, &marshalErr) && msg.Kind == jsonrpc.MessageKindResponse { + if resp := msg.AsResponse(); resp.ID != nil && resp.Error == nil { + s.logger.Errorf("failed to marshal response for request %s: %v", resp.ID, marshalErr) + if sendErr := s.sendError(resp.ID, marshalErr); sendErr != nil { + return sendErr + } + continue + } + } return fmt.Errorf("failed to write message: %w", err) } } diff --git a/internal/lsp/server_shutdown_test.go b/internal/lsp/server_test.go similarity index 50% rename from internal/lsp/server_shutdown_test.go rename to internal/lsp/server_test.go index f1dbc97bb38..6602f62fc01 100644 --- a/internal/lsp/server_shutdown_test.go +++ b/internal/lsp/server_test.go @@ -4,8 +4,10 @@ import ( "context" "io" "testing" + "time" "github.com/microsoft/typescript-go/internal/bundled" + "github.com/microsoft/typescript-go/internal/jsonrpc" "github.com/microsoft/typescript-go/internal/lsp/lsproto" "github.com/microsoft/typescript-go/internal/project" "github.com/microsoft/typescript-go/internal/vfs/vfstest" @@ -125,3 +127,108 @@ func TestServerOutgoingQueueDoesNotBlockWithoutWriter(t *testing.T) { t.Fatal("sending outgoing messages blocked without a writer") } } + +// A response that exceeds the JSON encoder's nesting limit must fail only its +// request. The write loop must remain available to deliver subsequent responses. +func TestWriteLoopRecoversFromUnserializableResponse(t *testing.T) { + t.Parallel() + + pr, pw := io.Pipe() + server := NewServer(&ServerOptions{ + In: shutdownTestReader{}, + Out: ToWriter(pw), + Err: io.Discard, + Cwd: "/test", + }) + + ctx, cancel := context.WithCancel(t.Context()) + defer cancel() + server.backgroundCtx = ctx + + writeLoopErr := make(chan error, 1) + go func() { writeLoopErr <- server.writeLoop(ctx) }() + + // A selection range whose parent chain is far deeper than the JSON encoder's nesting limit. + var deep *lsproto.SelectionRange + for range 20000 { + deep = &lsproto.SelectionRange{Parent: deep} + } + badResult := []*lsproto.SelectionRange{deep} + badID := jsonrpc.NewIDString("bad") + if err := server.send((&lsproto.ResponseMessage{ID: badID, Result: &badResult}).Message()); err != nil { + t.Fatalf("failed to enqueue bad response: %v", err) + } + + // A subsequent well-formed response must still be delivered. + goodID := jsonrpc.NewIDString("good") + if err := server.send((&lsproto.ResponseMessage{ID: goodID, Result: &lsproto.SelectionRangesOrNull{}}).Message()); err != nil { + t.Fatalf("failed to enqueue good response: %v", err) + } + + reader := lsproto.NewBaseReader(pr) + sawError := false + sawGood := false + for range 2 { + msg := readMessageWithTimeout(t, reader) + resp := msg.AsResponse() + switch { + case resp.ID != nil && *resp.ID == *badID: + if resp.Error == nil { + t.Errorf("expected an error response for the unserializable request, got a result") + } else if resp.Error.Code != int32(lsproto.ErrorCodeInternalError) { + t.Errorf("error response code = %d, want %d", resp.Error.Code, lsproto.ErrorCodeInternalError) + } + sawError = true + case resp.ID != nil && *resp.ID == *goodID: + if resp.Error != nil { + t.Errorf("expected a successful response for the good request, got error: %v", resp.Error) + } + sawGood = true + default: + t.Errorf("unexpected response id: %v", resp.ID) + } + } + + if !sawError { + t.Errorf("did not receive an error response for the unserializable request") + } + if !sawGood { + t.Errorf("did not receive the subsequent well-formed response (write loop likely died)") + } + + // The write loop must still be running. + select { + case err := <-writeLoopErr: + t.Fatalf("write loop exited unexpectedly: %v", err) + default: + return + } +} + +func readMessageWithTimeout(t *testing.T, reader *lsproto.BaseReader) *lsproto.Message { + t.Helper() + type result struct { + msg *lsproto.Message + err error + } + ch := make(chan result, 1) + go func() { + data, err := reader.Read() + if err != nil { + ch <- result{err: err} + return + } + msg := &lsproto.Message{} + ch <- result{msg: msg, err: msg.UnmarshalJSON(data)} + }() + select { + case r := <-ch: + if r.err != nil { + t.Fatalf("failed to read message: %v", r.err) + } + return r.msg + case <-time.After(2 * time.Second): + t.Fatal("timed out waiting for a message (write loop may have died)") + return nil + } +}