diff --git a/tsc/internal/ipc/conn_async.go b/tsc/internal/ipc/conn_async.go index 6094f8c1cd372..48b6bfc074cbd 100644 --- a/tsc/internal/ipc/conn_async.go +++ b/tsc/internal/ipc/conn_async.go @@ -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() @@ -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) } @@ -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) @@ -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 { @@ -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() @@ -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 @@ -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) } } }() @@ -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. diff --git a/tsc/internal/ipc/conn_async_test.go b/tsc/internal/ipc/conn_async_test.go index f5269216a2383..89a736e5fb1d8 100644 --- a/tsc/internal/ipc/conn_async_test.go +++ b/tsc/internal/ipc/conn_async_test.go @@ -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" ) @@ -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() @@ -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") + } +}