Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
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
56 changes: 44 additions & 12 deletions tsc/internal/ipc/conn_async.go
Original file line number Diff line number Diff line change
Expand Up @@ -64,7 +64,26 @@ func (c *AsyncConn) SetCollectTiming(enabled bool) {
// Run starts processing messages on the connection.
// It blocks until the context is cancelled or an error occurs.
func (c *AsyncConn) Run(ctx context.Context) (err error) {
defer func() { c.closePendingCalls(err) }()
ctx, cancel := context.WithCancel(ctx)
requestErrors := make(chan error, 1)
reportRequestError := func(requestErr error) {
select {
case requestErrors <- requestErr:
return
default:
return
}
}
defer func() {
cancel()
select {
case requestErr := <-requestErrors:
err = errors.Join(err, requestErr)
default:
// No request failed before the read loop exited.
}
c.closePendingCalls(err)
}()
for {
if ctx.Err() != nil {
return ctx.Err()
Expand All @@ -81,7 +100,12 @@ func (c *AsyncConn) Run(ctx context.Context) (err error) {
if msg.IsResponse() {
c.handleResponse(msg)
} else if msg.IsRequest() {
go c.handleRequest(ctx, msg)
go func() {
if requestErr := c.handleRequest(ctx, msg); requestErr != nil {
reportRequestError(requestErr)
_ = c.rwc.Close()
}
}()
} else if msg.IsNotification() {
go c.handleNotification(ctx, msg)
}
Expand All @@ -94,9 +118,9 @@ func (c *AsyncConn) closePendingCalls(runErr error) {
defer c.pendingMu.Unlock()
if c.terminal == nil {
c.terminal = ErrConnClosed
if runErr != nil {
c.terminal = errors.Join(c.terminal, runErr)
}
}
if runErr != nil {
c.terminal = errors.Join(c.terminal, runErr)
}
for id, ch := range c.pending {
close(ch)
Expand All @@ -120,7 +144,7 @@ func (c *AsyncConn) handleResponse(msg *Message) {
}

// handleRequest processes an incoming request.
func (c *AsyncConn) handleRequest(ctx context.Context, msg *Message) {
func (c *AsyncConn) handleRequest(ctx context.Context, msg *Message) (retErr error) {
// Intercept the meta-requests for collected server timing before dispatching
// to the handler, so they are answered directly and not themselves recorded.
switch msg.Method {
Expand All @@ -129,9 +153,11 @@ func (c *AsyncConn) handleRequest(ctx context.Context, msg *Message) {
writeErr := c.protocol.WriteResponse(msg.ID, serverTimingSnapshot(c.timing))
c.writeMu.Unlock()
if writeErr != nil {
panic(fmt.Sprintf("ipc: failed to write server timing response: %v", writeErr))
requestErr := fmt.Errorf("ipc: failed to write server timing response: %w", writeErr)
c.closePendingCalls(requestErr)
return requestErr
}
return
return nil
case string(MethodResetServerTiming):
if c.timing != nil {
c.timing.reset()
Expand All @@ -140,9 +166,11 @@ func (c *AsyncConn) handleRequest(ctx context.Context, msg *Message) {
writeErr := c.protocol.WriteResponse(msg.ID, nil)
c.writeMu.Unlock()
if writeErr != nil {
panic(fmt.Sprintf("ipc: failed to write reset server timing response: %v", writeErr))
requestErr := fmt.Errorf("ipc: failed to write reset server timing response: %w", writeErr)
c.closePendingCalls(requestErr)
return requestErr
}
return
return nil
}

var result any
Expand All @@ -167,7 +195,8 @@ func (c *AsyncConn) handleRequest(ctx context.Context, msg *Message) {
c.writeMu.Unlock()

if writeErr != nil {
panic(fmt.Sprintf("ipc: failed to write panic error response: %v (original panic: %v)", writeErr, r))
retErr = fmt.Errorf("ipc: failed to write panic error response: %w (original panic: %v)", writeErr, r)
c.closePendingCalls(retErr)
}
}
}()
Expand All @@ -192,8 +221,11 @@ func (c *AsyncConn) handleRequest(ctx context.Context, msg *Message) {
}

if writeErr != nil {
panic(fmt.Sprintf("ipc: failed to write response: %v", writeErr))
requestErr := fmt.Errorf("ipc: failed to write response: %w", writeErr)
c.closePendingCalls(requestErr)
return requestErr
}
return nil
}

// handleNotification processes an incoming notification.
Expand Down
146 changes: 146 additions & 0 deletions tsc/internal/ipc/conn_async_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -5,11 +5,13 @@ import (
"errors"
"io"
"net"
"sync"
"testing"
"time"

"github.com/microsoft/TypeScript/tsc/internal/ipc"
"github.com/microsoft/TypeScript/tsc/internal/json"
"github.com/microsoft/TypeScript/tsc/internal/jsonrpc"
"gotest.tools/v3/assert"
)

Expand All @@ -23,6 +25,81 @@ func (noOpHandler) HandleNotification(context.Context, string, json.Value) error
return nil
}

type blockingHandler struct {
started chan struct{}
release chan struct{}
}

func (h blockingHandler) HandleRequest(context.Context, string, json.Value) (any, error) {
close(h.started)
<-h.release
return nil, nil
}

func (blockingHandler) HandleNotification(context.Context, string, json.Value) error {
return nil
}

type closeSignal struct {
closed chan struct{}
once sync.Once
}

func (*closeSignal) Read([]byte) (int, error) {
return 0, io.EOF
}

func (*closeSignal) Write(p []byte) (int, error) {
return len(p), nil
}

func (c *closeSignal) Close() error {
c.once.Do(func() { close(c.closed) })
return nil
}

type failingResponseProtocol struct {
closed <-chan struct{}
requestRead bool
responseErr error
}

func (p *failingResponseProtocol) ReadMessage() (*ipc.Message, error) {
if !p.requestRead {
p.requestRead = true
return &ipc.Message{ID: jsonrpc.NewIDInt(1), Method: "transform"}, nil
}
<-p.closed
return nil, io.ErrClosedPipe
}

func (*failingResponseProtocol) WriteRequest(*jsonrpc.ID, string, any) error {
return nil
}

func (*failingResponseProtocol) WriteNotification(string, any) error {
return nil
}

func (p *failingResponseProtocol) WriteResponse(*jsonrpc.ID, any) error {
return p.responseErr
}

func (p *failingResponseProtocol) WriteError(*jsonrpc.ID, *jsonrpc.ResponseError) error {
return p.responseErr
}

type closeNotifyingReadWriteCloser struct {
io.ReadWriteCloser
closed chan struct{}
once sync.Once
}

func (c *closeNotifyingReadWriteCloser) Close() error {
c.once.Do(func() { close(c.closed) })
return c.ReadWriteCloser.Close()
}

func TestAsyncConnCallReturnsWhenPeerCloses(t *testing.T) {
t.Parallel()
client, server := net.Pipe()
Expand Down Expand Up @@ -67,3 +144,72 @@ func TestAsyncConnCallAfterReadLoopFailureReturnsImmediately(t *testing.T) {
err = conn.Notify(ctx, "changed", nil)
assert.Assert(t, errors.Is(err, ipc.ErrConnClosed), "expected ErrConnClosed, got %v", err)
}

func TestAsyncConnTerminalErrorIncludesResponseWriteFailure(t *testing.T) {
t.Parallel()
responseErr := errors.New("response write failed")
rwc := &closeSignal{closed: make(chan struct{})}
protocol := &failingResponseProtocol{
closed: rwc.closed,
responseErr: responseErr,
}
conn := ipc.NewAsyncConnWithProtocol(rwc, protocol, noOpHandler{})

err := conn.Run(t.Context())
assert.Assert(t, errors.Is(err, responseErr), "expected response write error, got %v", err)
_, err = conn.Call(t.Context(), "transform", nil)
assert.Assert(t, errors.Is(err, responseErr), "expected terminal response write error, got %v", err)
}

func TestAsyncConnRunReturnsWhenPeerClosesDuringRequest(t *testing.T) {
t.Parallel()
client, server := net.Pipe()
defer server.Close()
handler := blockingHandler{
started: make(chan struct{}),
release: make(chan struct{}),
}
defer func() {
select {
case <-handler.release:
return
default:
close(handler.release)
}
}()
serverTransport := &closeNotifyingReadWriteCloser{
ReadWriteCloser: server,
closed: make(chan struct{}),
}
conn := ipc.NewAsyncConn(serverTransport, handler)
runDone := make(chan error, 1)
go func() { runDone <- conn.Run(t.Context()) }()

clientProtocol := ipc.NewJSONRPCProtocol(client)
assert.NilError(t, clientProtocol.WriteRequest(jsonrpc.NewIDInt(1), "transform", nil))
select {
case <-handler.started:
break
case <-time.After(time.Second):
t.Fatal("request handler did not start")
}
assert.NilError(t, client.Close())

select {
case <-runDone:
break
case <-time.After(time.Second):
t.Fatal("connection did not stop while request handler was blocked")
}

close(handler.release)
select {
case <-serverTransport.closed:
_, err := conn.Call(t.Context(), "transform", nil)
assert.ErrorContains(t, err, "ipc: failed to write response")
err = conn.Notify(t.Context(), "changed", nil)
assert.ErrorContains(t, err, "ipc: failed to write response")
case <-time.After(time.Second):
t.Fatal("connection did not close after response write failure")
}
}