Skip to content

Commit 43481da

Browse files
committed
Make reaper handshake respect context deadlines
1 parent f18430b commit 43481da

2 files changed

Lines changed: 50 additions & 0 deletions

File tree

reaper.go

Lines changed: 8 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -542,6 +542,14 @@ func (r *Reaper) connect(ctx context.Context) (chan bool, error) {
542542
return nil, fmt.Errorf("dial reaper %s: %w", r.Endpoint, err)
543543
}
544544

545+
if deadline, ok := ctx.Deadline(); ok {
546+
if err := conn.SetDeadline(deadline); err != nil {
547+
conn.Close()
548+
return nil, fmt.Errorf("set handshake deadline for reaper %s: %w", r.Endpoint, err)
549+
}
550+
defer conn.SetDeadline(time.Time{})
551+
}
552+
545553
if err := r.handshake(conn); err != nil {
546554
conn.Close()
547555
return nil, fmt.Errorf("handshake reaper %s: %w", r.Endpoint, err)

reaper_test.go

Lines changed: 42 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -571,6 +571,48 @@ func TestReaperConnectReturnsHandshakeError(t *testing.T) {
571571
}
572572
}
573573

574+
func TestReaperConnectCancelsHungHandshake(t *testing.T) {
575+
ctx, cancel := context.WithTimeout(context.Background(), 500*time.Millisecond)
576+
defer cancel()
577+
578+
listener, err := net.Listen("tcp", "127.0.0.1:0")
579+
require.NoError(t, err)
580+
t.Cleanup(func() {
581+
require.NoError(t, listener.Close())
582+
})
583+
584+
done := make(chan struct{})
585+
go func() {
586+
defer close(done)
587+
588+
conn, err := listener.Accept()
589+
if err != nil {
590+
return
591+
}
592+
defer conn.Close()
593+
594+
_, _ = bufio.NewReader(conn).ReadString('\n')
595+
<-ctx.Done()
596+
}()
597+
598+
reaper := &Reaper{
599+
SessionID: testSessionID,
600+
Endpoint: listener.Addr().String(),
601+
}
602+
603+
termSignal, err := reaper.connect(ctx)
604+
require.Nil(t, termSignal)
605+
require.ErrorContains(t, err, "handshake reaper")
606+
require.ErrorContains(t, err, "read ack")
607+
require.ErrorIs(t, err, os.ErrDeadlineExceeded)
608+
609+
select {
610+
case <-done:
611+
case <-time.After(time.Second):
612+
require.FailNow(t, "test reaper server did not finish")
613+
}
614+
}
615+
574616
func reaperConnect(t *testing.T, reaper *Reaper) {
575617
t.Helper()
576618

0 commit comments

Comments
 (0)