Skip to content
Merged
Show file tree
Hide file tree
Changes from 1 commit
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
20 changes: 19 additions & 1 deletion internal/lsp/server.go
Original file line number Diff line number Diff line change
Expand Up @@ -113,6 +113,14 @@ type lspWriter struct {
w *lsproto.BaseWriter
}

// MessageMarshalError indicates that an outgoing message could not be serialized to JSON
type MessageMarshalError struct {
Comment thread
johnfav03 marked this conversation as resolved.
Outdated
err error
}

func (e *MessageMarshalError) Error() string { return "failed to marshal message: " + e.err.Error() }
func (e *MessageMarshalError) Unwrap() error { return e.err }

func (r *lspReader) Read() (*lsproto.Message, error) {
data, err := r.r.Read()
if err != nil {
Expand All @@ -137,7 +145,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)
}
Expand Down Expand Up @@ -626,6 +634,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)
}
}
Expand Down
118 changes: 118 additions & 0 deletions internal/lsp/server_marshal_failure_test.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,118 @@
package lsp
Comment thread
johnfav03 marked this conversation as resolved.
Outdated

import (
"context"
"io"
"testing"
"time"

"github.com/microsoft/typescript-go/internal/jsonrpc"
"github.com/microsoft/typescript-go/internal/lsp/lsproto"
)

type eofReader struct{}

func (eofReader) Read() (*lsproto.Message, error) { return nil, io.EOF }

// 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: eofReader{},
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 unmarshalable request, got a result")
}
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 unmarshalable 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
}
}