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
7 changes: 7 additions & 0 deletions CHANGELOG.md
Original file line number Diff line number Diff line change
Expand Up @@ -3,6 +3,13 @@ Changelog

All notable changes to this project will be documented in this file.

## Unreleased

### Fixed

- Input `websocket`: Connection attempts (dial and upgrade handshake) now honor context cancellation. Previously, unresponsive servers could
block graceful shutdown for up to the default 45s handshake timeout. (@Leward)

## 4.78.0 - 2026-08-20

### Added
Expand Down
8 changes: 3 additions & 5 deletions internal/impl/io/input_websocket.go
Original file line number Diff line number Diff line change
Expand Up @@ -185,13 +185,11 @@ func (w *websocketReader) getConn(ctx context.Context) (*websocket.Conn, error)
if w.proxyURLParsed != nil {
dialer.Proxy = http.ProxyURL(w.proxyURLParsed)
}

if w.tlsEnabled {
dialer.TLSClientConfig = w.tlsConf
if client, res, err = dialer.Dial(w.urlStr, headers); err != nil {
return nil, err
}
} else if client, res, err = dialer.Dial(w.urlStr, headers); err != nil {

Copy link
Copy Markdown
Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

The branching was redundant and has been simplified

}

if client, res, err = dialContext(ctx, dialer, w.urlStr, headers); err != nil {
return nil, err
}

Expand Down
106 changes: 106 additions & 0 deletions internal/impl/io/input_websocket_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -6,6 +6,7 @@ import (
"bytes"
"context"
"fmt"
"net"
"net/http"
"net/http/httptest"
"net/url"
Expand Down Expand Up @@ -304,3 +305,108 @@ url: %v
wg.Wait()
close(closeChan)
}

// newWebsocketReader returns a websocketReader configured to connect to the websocket server listening on addr.
func newWebsocketReader(t *testing.T, addr net.Addr) *websocketReader {
t.Helper()

pConf, err := websocketInputSpec().ParseYAML("url: ws://"+addr.String()+"\n", nil)
require.NoError(t, err)

m, err := newWebsocketReaderFromParsed(pConf, mock.NewManager())
require.NoError(t, err)

return m
}

// newHangingWebsocketReader returns a reader pointed at a listener that accepts
// TCP connections but never answers the handshake, along with the accept channel
// of that listener.
func newHangingWebsocketReader(t *testing.T) (*websocketReader, <-chan struct{}) {
t.Helper()

addr, accepted := newHangingListener(t)
return newWebsocketReader(t, addr), accepted
}

// newUnreachableWebsocketReader returns a reader pointed at a closed port. A dial
// there fails immediately, so a context error proves no dial was attempted.
func newUnreachableWebsocketReader(t *testing.T) *websocketReader {
t.Helper()

return newWebsocketReader(t, newUnreachableAddr(t))
}

// TestWebsocketConnectContextDone tests that Connect reports the context error
// when the context is done before or during the websocket handshake.
func TestWebsocketConnectContextDone(t *testing.T) {
tests := []struct {
name string
// setUp returns a reader and a context that is already done, or that becomes
// done while the dial waits for the handshake response. It also returns the
// accept channel to check after Connect returns, or nil for a case where no
// accept is expected.
setUp func(*testing.T) (*websocketReader, context.Context, <-chan struct{})
wantErr error
}{
{
name: "canceled before dialing",
setUp: func(t *testing.T) (*websocketReader, context.Context, <-chan struct{}) {
ctx, cancel := context.WithCancel(t.Context())
cancel()
// The port is closed, so a dial fails at once. Only a context check
// before the dial can give context.Canceled here.
return newUnreachableWebsocketReader(t), ctx, nil
},
wantErr: context.Canceled,
},
{
name: "canceled while dialing",
setUp: func(t *testing.T) (*websocketReader, context.Context, <-chan struct{}) {
m, accepted := newHangingWebsocketReader(t)

ctx, cancel := context.WithCancel(t.Context())
t.Cleanup(cancel)
// Cancel from the accept, not from a timer, so the dial always waits
// for the handshake response when the context becomes done.
go func() {
<-accepted
cancel()
}()

return m, ctx, nil
},
wantErr: context.Canceled,
},
{
name: "deadline exceeded while dialing",
setUp: func(t *testing.T) (*websocketReader, context.Context, <-chan struct{}) {
m, accepted := newHangingWebsocketReader(t)

// A context carries its deadline from the start, so this case cannot
// take its trigger from the accept. The accept check after Connect
// returns proves the deadline expired during the handshake.
ctx, cancel := context.WithTimeout(t.Context(), 50*time.Millisecond)
t.Cleanup(cancel)

return m, ctx, accepted
},
wantErr: context.DeadlineExceeded,
},
}

for _, test := range tests {
t.Run(test.name, func(t *testing.T) {
m, ctx, accepted := test.setUp(t)

err := awaitWithin(t, 500*time.Millisecond, "Connect", func() error {
return m.Connect(ctx)
})
require.ErrorIs(t, err, test.wantErr)

if accepted != nil {
requireAccepted(t, 5*time.Second, accepted)
}
})
}
}
69 changes: 69 additions & 0 deletions internal/impl/io/websocket.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,69 @@
// Copyright 2026 Redpanda Data, Inc.

package io

import (
"context"
"net"
"net/http"
"time"

"github.com/gorilla/websocket"
)

// hiddenDeadlineContext strips ctx.Deadline() while preserving cancellation.
//
// Gorilla sets a socket deadline matching ctx.Deadline(). Because the runtime's
// socket deadline timer and Go's context timer run independently, a socket read
// can time out (returning a generic "i/o timeout") slightly before ctx.Err() is set.
//
// Concealing the deadline prevents Gorilla from setting this competing socket timer.
// This ensures our context.AfterFunc watcher exclusively controls socket interruption
// and deterministically returns context errors.
type hiddenDeadlineContext struct {
context.Context
}

func (hiddenDeadlineContext) Deadline() (time.Time, bool) {
return time.Time{}, false
}

// dialContext is dialer.DialContext plus cancellation of the HTTP upgrade exchange.
//
// Gorilla applies the context to the TCP and TLS handshakes only.
// It then performs the websocket upgrade on a bare connection, which does not respond to context cancellation.
// A peer that accepts TCP and then stays silent blocks the dial until the handshake timeout.
func dialContext(ctx context.Context, dialer websocket.Dialer, urlStr string, headers http.Header) (*websocket.Conn, *http.Response, error) {
netDialer := &net.Dialer{}
var stop func() bool

// 1. Intercept network connection creation to attach a context watcher (context.AfterFunc) that closes the conn once ctx is done.
dialer.NetDialContext = func(dialCtx context.Context, network, addr string) (net.Conn, error) {
conn, err := netDialer.DialContext(dialCtx, network, addr)
if err != nil {
return nil, err
}
// Watch ctx, not dialCtx: gorilla wraps dialCtx with HandshakeTimeout and
// cancels it as DialContext returns, which would fire this on success.
stop = context.AfterFunc(ctx, func() { _ = conn.Close() })
return conn, nil
}

// 2. Perform dial with hidden deadline context.
// The context we give gorilla keeps the cancellation of ctx (stopping TCP/TLS
// handshakes), but conceals ctx.Deadline() so gorilla does not set a competing
// socket timer that could race with ctx cancellation.
client, res, err := dialer.DialContext(hiddenDeadlineContext{ctx}, urlStr, headers)

// 3. Clean up and prioritize context error returns.
if stop != nil {
_ = stop()
}
if ctxErr := ctx.Err(); ctxErr != nil {
if client != nil {
_ = client.Close()
}
return nil, res, ctxErr
}
return client, res, err
}
103 changes: 103 additions & 0 deletions internal/impl/io/websocket_test.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,103 @@
// Copyright 2026 Redpanda Data, Inc.

package io

import (
"net"
"sync"
"testing"
"time"

"github.com/stretchr/testify/require"
)

// newHangingListener returns the address of a listener that accepts TCP
// connections but never answers the handshake, so a dial there stays blocked
// until the dialer times out. The returned channel reports each accept, which
// tells the caller that a handshake is in flight.
func newHangingListener(t *testing.T) (net.Addr, <-chan struct{}) {
t.Helper()

lis, err := net.Listen("tcp", "127.0.0.1:0")
require.NoError(t, err)

accepted := make(chan struct{}, 1)
var (
mut sync.Mutex
conns []net.Conn
)

go func() {
for {
conn, err := lis.Accept()
if err != nil {
return
}
mut.Lock()
conns = append(conns, conn)
mut.Unlock()
select {
case accepted <- struct{}{}:
default:
}
}
}()

// Close the accepted connections as well as the listener, so that an
// abandoned dial fails instead of running to the handshake timeout.
t.Cleanup(func() {
_ = lis.Close()
mut.Lock()
defer mut.Unlock()
for _, conn := range conns {
_ = conn.Close()
}
})

return lis.Addr(), accepted
}

// newUnreachableAddr returns the address of a closed port. A dial there fails
// immediately, so a context error proves no dial was attempted.
func newUnreachableAddr(t *testing.T) net.Addr {
t.Helper()

lis, err := net.Listen("tcp", "127.0.0.1:0")
require.NoError(t, err)
addr := lis.Addr()
require.NoError(t, lis.Close())

return addr
}

// awaitWithin runs fn in the background and returns its error, or fails the test
// if fn has not returned within d.
func awaitWithin(t *testing.T, d time.Duration, what string, fn func() error) error {
t.Helper()

done := make(chan error, 1)
go func() {
done <- fn()
}()

select {
case err := <-done:
return err
case <-time.After(d):
t.Fatalf("%v stayed blocked for %v", what, d)
return nil
}
}

// requireAccepted fails the test if the listener reported no accept within d. It
// proves that the dial reached the handshake, so a test that expects the context
// to become mid-handshake cannot pass through an earlier failure instead.
func requireAccepted(t *testing.T, d time.Duration, accepted <-chan struct{}) {
t.Helper()

select {
case <-accepted:
case <-time.After(d):
t.Fatal("the listener accepted no connection, so the dial never reached the handshake")
}
}
Loading