Skip to content

Commit 89600d9

Browse files
committed
Handle IPC response write failures
Return response write failures through sync and async connections instead of panicking while handling requests. Preserve the first async terminal cause, unblock pending calls, and close the transport once while retaining handler teardown ordering.
1 parent 725f612 commit 89600d9

4 files changed

Lines changed: 281 additions & 21 deletions

File tree

‎tsc/internal/ipc/conn_async.go‎

Lines changed: 51 additions & 10 deletions
Original file line numberDiff line numberDiff line change
@@ -31,6 +31,7 @@ type AsyncConn struct {
3131
pending map[jsonrpc.ID]chan *Message
3232
pendingMu sync.Mutex
3333
terminal error
34+
hasCause bool
3435
writeMu sync.Mutex
3536
handlers sync.WaitGroup
3637
}
@@ -66,10 +67,17 @@ func (c *AsyncConn) SetCollectTiming(enabled bool) {
6667
// It blocks until the context is cancelled or an error occurs.
6768
func (c *AsyncConn) Run(ctx context.Context) (err error) {
6869
handlerCtx, cancelHandlers := context.WithCancel(ctx)
70+
requestErrors := make(chan error, 1)
6971
defer func() {
7072
c.closePendingCalls(err)
7173
cancelHandlers()
7274
c.handlers.Wait()
75+
select {
76+
case requestErr := <-requestErrors:
77+
err = errors.Join(err, requestErr)
78+
default:
79+
// No request failed before the read loop exited.
80+
}
7381
}()
7482
for {
7583
if ctx.Err() != nil {
@@ -88,7 +96,13 @@ func (c *AsyncConn) Run(ctx context.Context) (err error) {
8896
c.handleResponse(msg)
8997
} else if msg.IsRequest() {
9098
c.handlers.Go(func() {
91-
c.handleRequest(handlerCtx, msg)
99+
if requestErr := c.handleRequest(handlerCtx, msg); requestErr != nil {
100+
if c.recordRequestError(requestErr, requestErrors) {
101+
if c.rwc != nil {
102+
_ = c.rwc.Close()
103+
}
104+
}
105+
}
92106
})
93107
} else if msg.IsNotification() {
94108
c.handlers.Go(func() {
@@ -102,12 +116,38 @@ func (c *AsyncConn) Run(ctx context.Context) (err error) {
102116
func (c *AsyncConn) closePendingCalls(runErr error) {
103117
c.pendingMu.Lock()
104118
defer c.pendingMu.Unlock()
119+
c.recordTerminalErrorLocked(runErr)
120+
c.closePendingCallsLocked()
121+
}
122+
123+
func (c *AsyncConn) recordRequestError(requestErr error, requestErrors chan<- error) bool {
124+
c.pendingMu.Lock()
125+
defer c.pendingMu.Unlock()
126+
if !c.recordTerminalErrorLocked(requestErr) {
127+
return false
128+
}
129+
requestErrors <- requestErr
130+
c.closePendingCallsLocked()
131+
return true
132+
}
133+
134+
func (c *AsyncConn) recordTerminalErrorLocked(terminalErr error) bool {
105135
if c.terminal == nil {
106136
c.terminal = ErrConnClosed
107-
if runErr != nil {
108-
c.terminal = errors.Join(c.terminal, runErr)
137+
if terminalErr != nil {
138+
c.terminal = errors.Join(c.terminal, terminalErr)
139+
c.hasCause = true
140+
return true
109141
}
142+
} else if !c.hasCause && terminalErr != nil {
143+
c.terminal = errors.Join(c.terminal, terminalErr)
144+
c.hasCause = true
145+
return true
110146
}
147+
return false
148+
}
149+
150+
func (c *AsyncConn) closePendingCallsLocked() {
111151
for id, ch := range c.pending {
112152
close(ch)
113153
delete(c.pending, id)
@@ -130,7 +170,7 @@ func (c *AsyncConn) handleResponse(msg *Message) {
130170
}
131171

132172
// handleRequest processes an incoming request.
133-
func (c *AsyncConn) handleRequest(ctx context.Context, msg *Message) {
173+
func (c *AsyncConn) handleRequest(ctx context.Context, msg *Message) (retErr error) {
134174
// Intercept the meta-requests for collected server timing before dispatching
135175
// to the handler, so they are answered directly and not themselves recorded.
136176
switch msg.Method {
@@ -139,9 +179,9 @@ func (c *AsyncConn) handleRequest(ctx context.Context, msg *Message) {
139179
writeErr := c.protocol.WriteResponse(msg.ID, serverTimingSnapshot(c.timing))
140180
c.writeMu.Unlock()
141181
if writeErr != nil {
142-
panic(fmt.Sprintf("ipc: failed to write server timing response: %v", writeErr))
182+
return fmt.Errorf("ipc: failed to write server timing response: %w", writeErr)
143183
}
144-
return
184+
return nil
145185
case string(MethodResetServerTiming):
146186
if c.timing != nil {
147187
c.timing.reset()
@@ -150,9 +190,9 @@ func (c *AsyncConn) handleRequest(ctx context.Context, msg *Message) {
150190
writeErr := c.protocol.WriteResponse(msg.ID, nil)
151191
c.writeMu.Unlock()
152192
if writeErr != nil {
153-
panic(fmt.Sprintf("ipc: failed to write reset server timing response: %v", writeErr))
193+
return fmt.Errorf("ipc: failed to write reset server timing response: %w", writeErr)
154194
}
155-
return
195+
return nil
156196
}
157197

158198
var result any
@@ -177,7 +217,7 @@ func (c *AsyncConn) handleRequest(ctx context.Context, msg *Message) {
177217
c.writeMu.Unlock()
178218

179219
if writeErr != nil {
180-
panic(fmt.Sprintf("ipc: failed to write panic error response: %v (original panic: %v)", writeErr, r))
220+
retErr = fmt.Errorf("ipc: failed to write panic error response: %w (original panic: %v)", writeErr, r)
181221
}
182222
}
183223
}()
@@ -202,8 +242,9 @@ func (c *AsyncConn) handleRequest(ctx context.Context, msg *Message) {
202242
}
203243

204244
if writeErr != nil {
205-
panic(fmt.Sprintf("ipc: failed to write response: %v", writeErr))
245+
return fmt.Errorf("ipc: failed to write response: %w", writeErr)
206246
}
247+
return nil
207248
}
208249

209250
// handleNotification processes an incoming notification.

‎tsc/internal/ipc/conn_async_test.go‎

Lines changed: 139 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -5,6 +5,8 @@ import (
55
"errors"
66
"io"
77
"net"
8+
"strings"
9+
"sync"
810
"testing"
911
"time"
1012

@@ -25,7 +27,8 @@ func (noOpHandler) HandleNotification(context.Context, string, json.Value) error
2527
}
2628

2729
type queuedProtocol struct {
28-
messages []*ipc.Message
30+
messages []*ipc.Message
31+
responseErr error
2932
}
3033

3134
func (p *queuedProtocol) ReadMessage() (*ipc.Message, error) {
@@ -46,11 +49,11 @@ func (p *queuedProtocol) WriteNotification(string, any) error {
4649
}
4750

4851
func (p *queuedProtocol) WriteResponse(*jsonrpc.ID, any) error {
49-
return nil
52+
return p.responseErr
5053
}
5154

5255
func (p *queuedProtocol) WriteError(*jsonrpc.ID, *jsonrpc.ResponseError) error {
53-
return nil
56+
return p.responseErr
5457
}
5558

5659
type blockingHandler struct {
@@ -132,6 +135,70 @@ func TestAsyncConnRunCancelsHandlersOnEOF(t *testing.T) {
132135
}
133136
}
134137

138+
func TestAsyncConnResponseWriteFailureWithNilTransport(t *testing.T) {
139+
t.Parallel()
140+
141+
responseErr := errors.New("response write failed")
142+
id := jsonrpc.NewIDString("1")
143+
protocol := &queuedProtocol{
144+
messages: []*ipc.Message{{ID: id, Method: "request"}},
145+
responseErr: responseErr,
146+
}
147+
conn := ipc.NewAsyncConnWithProtocol(nil, protocol, noOpHandler{})
148+
149+
err := conn.Run(t.Context())
150+
assert.Assert(t, errors.Is(err, responseErr), "expected response write error, got %v", err)
151+
}
152+
153+
type closeSignal struct {
154+
closed chan struct{}
155+
once sync.Once
156+
}
157+
158+
func (*closeSignal) Read([]byte) (int, error) {
159+
return 0, io.EOF
160+
}
161+
162+
func (*closeSignal) Write(p []byte) (int, error) {
163+
return len(p), nil
164+
}
165+
166+
func (c *closeSignal) Close() error {
167+
c.once.Do(func() { close(c.closed) })
168+
return nil
169+
}
170+
171+
type failingResponseProtocol struct {
172+
closed <-chan struct{}
173+
requestRead bool
174+
responseErr error
175+
}
176+
177+
func (p *failingResponseProtocol) ReadMessage() (*ipc.Message, error) {
178+
if !p.requestRead {
179+
p.requestRead = true
180+
return &ipc.Message{ID: jsonrpc.NewIDInt(1), Method: "transform"}, nil
181+
}
182+
<-p.closed
183+
return nil, io.ErrClosedPipe
184+
}
185+
186+
func (*failingResponseProtocol) WriteRequest(*jsonrpc.ID, string, any) error {
187+
return nil
188+
}
189+
190+
func (*failingResponseProtocol) WriteNotification(string, any) error {
191+
return nil
192+
}
193+
194+
func (p *failingResponseProtocol) WriteResponse(*jsonrpc.ID, any) error {
195+
return p.responseErr
196+
}
197+
198+
func (p *failingResponseProtocol) WriteError(*jsonrpc.ID, *jsonrpc.ResponseError) error {
199+
return p.responseErr
200+
}
201+
135202
func TestAsyncConnCallReturnsWhenPeerCloses(t *testing.T) {
136203
t.Parallel()
137204
client, server := net.Pipe()
@@ -176,3 +243,72 @@ func TestAsyncConnCallAfterReadLoopFailureReturnsImmediately(t *testing.T) {
176243
err = conn.Notify(ctx, "changed", nil)
177244
assert.Assert(t, errors.Is(err, ipc.ErrConnClosed), "expected ErrConnClosed, got %v", err)
178245
}
246+
247+
func TestAsyncConnTerminalErrorIncludesResponseWriteFailure(t *testing.T) {
248+
t.Parallel()
249+
responseErr := errors.New("response write failed")
250+
rwc := &closeSignal{closed: make(chan struct{})}
251+
protocol := &failingResponseProtocol{
252+
closed: rwc.closed,
253+
responseErr: responseErr,
254+
}
255+
conn := ipc.NewAsyncConnWithProtocol(rwc, protocol, noOpHandler{})
256+
257+
err := conn.Run(t.Context())
258+
assert.Assert(t, errors.Is(err, responseErr), "expected response write error, got %v", err)
259+
_, err = conn.Call(t.Context(), "transform", nil)
260+
assert.Assert(t, errors.Is(err, responseErr), "expected terminal response write error, got %v", err)
261+
assert.Equal(t, strings.Count(err.Error(), responseErr.Error()), 1)
262+
}
263+
264+
func TestAsyncConnRunWaitsForRequestAfterPeerCloses(t *testing.T) {
265+
t.Parallel()
266+
client, server := net.Pipe()
267+
defer server.Close()
268+
handler := &blockingHandler{
269+
started: make(chan struct{}, 1),
270+
release: make(chan struct{}),
271+
}
272+
defer func() {
273+
select {
274+
case <-handler.release:
275+
return
276+
default:
277+
close(handler.release)
278+
}
279+
}()
280+
conn := ipc.NewAsyncConn(server, handler)
281+
runDone := make(chan error, 1)
282+
go func() { runDone <- conn.Run(t.Context()) }()
283+
284+
clientProtocol := ipc.NewJSONRPCProtocol(client)
285+
assert.NilError(t, clientProtocol.WriteRequest(jsonrpc.NewIDInt(1), "transform", nil))
286+
select {
287+
case <-handler.started:
288+
break
289+
case <-time.After(time.Second):
290+
t.Fatal("request handler did not start")
291+
}
292+
assert.NilError(t, client.Close())
293+
294+
handlerBlocked := false
295+
select {
296+
case err := <-runDone:
297+
t.Fatalf("connection stopped while request handler was blocked: %v", err)
298+
case <-time.After(100 * time.Millisecond):
299+
handlerBlocked = true
300+
}
301+
assert.Assert(t, handlerBlocked)
302+
303+
close(handler.release)
304+
select {
305+
case err := <-runDone:
306+
assert.ErrorContains(t, err, "ipc: failed to write response")
307+
_, err = conn.Call(t.Context(), "transform", nil)
308+
assert.ErrorContains(t, err, "ipc: failed to write response")
309+
err = conn.Notify(t.Context(), "changed", nil)
310+
assert.ErrorContains(t, err, "ipc: failed to write response")
311+
case <-time.After(time.Second):
312+
t.Fatal("connection did not stop after request handler completed")
313+
}
314+
}

‎tsc/internal/ipc/conn_sync.go‎

Lines changed: 11 additions & 8 deletions
Original file line numberDiff line numberDiff line change
@@ -70,7 +70,9 @@ func (c *SyncConn) Run(ctx context.Context) error {
7070
}
7171

7272
if msg.IsRequest() {
73-
c.handleRequest(ctx, msg)
73+
if err := c.handleRequest(ctx, msg); err != nil {
74+
return err
75+
}
7476
} else if msg.IsNotification() {
7577
c.handleNotification(ctx, msg)
7678
} else {
@@ -81,7 +83,7 @@ func (c *SyncConn) Run(ctx context.Context) error {
8183
}
8284

8385
// handleRequest processes an incoming request.
84-
func (c *SyncConn) handleRequest(ctx context.Context, msg *Message) {
86+
func (c *SyncConn) handleRequest(ctx context.Context, msg *Message) (retErr error) {
8587
// Intercept the meta-requests for collected server timing before dispatching
8688
// to the handler, so they are answered directly and not themselves recorded.
8789
switch msg.Method {
@@ -90,9 +92,9 @@ func (c *SyncConn) handleRequest(ctx context.Context, msg *Message) {
9092
writeErr := c.protocol.WriteResponse(msg.ID, serverTimingSnapshot(c.timing))
9193
c.mu.Unlock()
9294
if writeErr != nil {
93-
panic(fmt.Sprintf("ipc: failed to write server timing response: %v", writeErr))
95+
return fmt.Errorf("ipc: failed to write server timing response: %w", writeErr)
9496
}
95-
return
97+
return nil
9698
case string(MethodResetServerTiming):
9799
if c.timing != nil {
98100
c.timing.reset()
@@ -101,9 +103,9 @@ func (c *SyncConn) handleRequest(ctx context.Context, msg *Message) {
101103
writeErr := c.protocol.WriteResponse(msg.ID, nil)
102104
c.mu.Unlock()
103105
if writeErr != nil {
104-
panic(fmt.Sprintf("ipc: failed to write reset server timing response: %v", writeErr))
106+
return fmt.Errorf("ipc: failed to write reset server timing response: %w", writeErr)
105107
}
106-
return
108+
return nil
107109
}
108110

109111
var result any
@@ -128,7 +130,7 @@ func (c *SyncConn) handleRequest(ctx context.Context, msg *Message) {
128130
c.mu.Unlock()
129131

130132
if writeErr != nil {
131-
panic(fmt.Sprintf("ipc: failed to write panic error response: %v (original panic: %v)", writeErr, r))
133+
retErr = fmt.Errorf("ipc: failed to write panic error response: %w (original panic: %v)", writeErr, r)
132134
}
133135
}
134136
}()
@@ -153,8 +155,9 @@ func (c *SyncConn) handleRequest(ctx context.Context, msg *Message) {
153155
}
154156

155157
if writeErr != nil {
156-
panic(fmt.Sprintf("ipc: failed to write response: %v", writeErr))
158+
return fmt.Errorf("ipc: failed to write response: %w", writeErr)
157159
}
160+
return nil
158161
}
159162

160163
// handleNotification processes an incoming notification.

0 commit comments

Comments
 (0)