Skip to content
Merged
4 changes: 4 additions & 0 deletions experimental/ssh/internal/proxy/client.go
Original file line number Diff line number Diff line change
Expand Up @@ -108,6 +108,10 @@ func RunClientProxy(ctx context.Context, src io.ReadCloser, dst io.Writer, reque
log.Debugf(gCtx, "Could not open a replacement connection for the auth handover, staying on the current one: %v", err)
continue
}
// The session ended during the handover: see the gCtx.Done case above.
if errors.Is(err, context.Canceled) {
return nil
}
if resumable {
// The failed handover closes its sockets. The receiving loop
// reattaches after the handover releases the write lock.
Expand Down
177 changes: 171 additions & 6 deletions experimental/ssh/internal/proxy/client_server_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -34,9 +34,11 @@ func createTestServer(t *testing.T, maxClients int, shutdownDelay time.Duration)
}

type testClient struct {
InputWriter io.WriteCloser
InputWriter *io.PipeWriter
Output *testBuffer
Cleanup func()
// Done is closed when RunClientProxy returns.
Done <-chan struct{}
Cleanup func()
}

func createTestClient(t *testing.T, serverURL string, requestHandoverTick func() <-chan time.Time, keepaliveInterval time.Duration, errChan chan error) *testClient {
Expand All @@ -58,8 +60,9 @@ func createTestClientWithDialer(t *testing.T, createConn createWebsocketConnecti
if requestHandoverTick == nil {
requestHandoverTick = neverTick
}
wg := sync.WaitGroup{}
wg.Go(func() {
done := make(chan struct{})
go func() {
defer close(done)
err := RunClientProxy(ctx, clientInput, clientOutput, requestHandoverTick, keepaliveInterval, resumable, createConn)
if err != nil && !errors.Is(err, context.Canceled) && !errors.Is(err, io.ErrClosedPipe) {
if errChan != nil {
Expand All @@ -68,14 +71,15 @@ func createTestClientWithDialer(t *testing.T, createConn createWebsocketConnecti
t.Errorf("client error: %v", err)
}
}
})
}()
return &testClient{
InputWriter: clientInputWriter,
Output: clientOutput,
Done: done,
Cleanup: func() {
clientInput.Close()
clientInputWriter.Close()
wg.Wait()
<-done
},
}
}
Expand Down Expand Up @@ -316,6 +320,167 @@ func TestHandoverDialFailureKeepsSessionAlive(t *testing.T) {
assert.Equal(t, int32(2), dials.Load(), "expected the initial dial plus exactly one handover dial")
}

// A failed dial never touches the live connection, whatever error it returns, so the handover loop
// must check for it before it treats context.Canceled as the session ending. Later ticks must still
// hand over.
func TestHandoverDialCanceledKeepsHandingOver(t *testing.T) {
server := createTestServer(t, 2, time.Hour)
defer server.Close()

wsURL := "ws" + server.URL[4:]
var dials atomic.Int32
handoverDialed := make(chan struct{}, 2)
createConn := func(ctx context.Context, dial DialRequest) (*websocket.Conn, error) {
n := dials.Add(1)
if n > 1 {
handoverDialed <- struct{}{}
}
if n == 2 {
return nil, context.Canceled
}
url := fmt.Sprintf("%s?id=%s", wsURL, dial.ConnID)
conn, _, err := websocket.DefaultDialer.Dial(url, nil) // nolint:bodyclose
return conn, err
}

handoverChan := make(chan time.Time)
client := createTestClientWithDialer(t, createConn, func() <-chan time.Time {
return handoverChan
}, time.Hour, false, nil)
defer client.Cleanup()

for i := range 2 {
select {
case handoverChan <- time.Now():
case <-time.After(10 * time.Second):
t.Fatalf("handover %d was never started: the handover loop has stopped", i+1)
}
// The tick only proves the loop received it. Wait for the dial, which runs under the
// handover mutex, so the write below cannot overtake the handover.
select {
case <-handoverDialed:
case <-time.After(10 * time.Second):
t.Fatalf("handover %d never dialed", i+1)
}
// sendMessage waits for the handover mutex, so this round trip also waits out the handover.
msg := fmt.Appendf(nil, "after handover %d\n", i+1)
_, err := client.InputWriter.Write(msg)
require.NoError(t, err)
require.NoError(t, client.Output.AssertWrite(msg))
}
assert.Equal(t, int32(3), dials.Load(), "expected the initial dial plus two handover dials")
}

// rawDrainServer accepts a websocket and reads its raw bytes without processing any frame, so it never
// answers anything, not even a close frame.
func rawDrainServer(t *testing.T) *httptest.Server {
return httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
conn, err := (&websocket.Upgrader{}).Upgrade(w, r, nil)
if err != nil {
return
}
defer conn.Close()
_, _ = io.Copy(io.Discard, conn.NetConn())
}))
}

// A session that ends while a handover is still waiting for the old connection to close must end
// with the session's own outcome, not ErrHandoverFailed, and must not start another handover.
func TestSessionEndDuringHandover(t *testing.T) {
errSourceFailed := errors.New("source failed")
tests := []struct {
name string
resumable bool
endSession func(w *io.PipeWriter)
wantErr error
}{
{
name: "clean exit",
endSession: func(w *io.PipeWriter) { w.Close() },
},
{
name: "source error",
endSession: func(w *io.PipeWriter) { w.CloseWithError(errSourceFailed) },
wantErr: errSourceFailed,
},
{
name: "resumable clean exit",
resumable: true,
endSession: func(w *io.PipeWriter) { w.Close() },
},
{
name: "resumable source error",
resumable: true,
endSession: func(w *io.PipeWriter) { w.CloseWithError(errSourceFailed) },
wantErr: errSourceFailed,
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
// Both connections land on a server that never answers, so nobody closes the old
// connection and the handover stays in progress until the session ends. It does not
// even echo the close frame of the client's teardown: an echo that arrived before
// teardown closed the socket would complete the swap instead.
server := rawDrainServer(t)
defer server.Close()

wsURL := "ws" + server.URL[4:]
var dials atomic.Int32
handoverDialed := make(chan struct{}, 1)
createConn := func(ctx context.Context, dial DialRequest) (*websocket.Conn, error) {
isHandover := dials.Add(1) > 1
conn, _, err := websocket.DefaultDialer.Dial(wsURL, nil) // nolint:bodyclose
// Only a successful dial leaves the handover in progress.
if isHandover && err == nil {
select {
case handoverDialed <- struct{}{}:
default:
}
}
return conn, err
}

// The second tick waits in the channel while the first handover is in progress, so a
// handover loop that keeps going after the session ended dials again.
ticks := make(chan time.Time, 2)
ticks <- time.Now()
ticks <- time.Now()
errChan := make(chan error, 1)
client := createTestClientWithDialer(t, createConn, func() <-chan time.Time {
return ticks
}, time.Hour, tt.resumable, errChan)
defer client.Cleanup()

select {
case <-handoverDialed:
case <-time.After(10 * time.Second):
t.Fatal("the handover never dialed its replacement connection")
}

// Wait for the session to end before Cleanup closes the reader too: a read that
// wakes up after that returns io.ErrClosedPipe instead of what endSession set.
tt.endSession(client.InputWriter)
select {
case <-client.Done:
case <-time.After(10 * time.Second):
t.Fatal("the session did not end")
}
var sessionErr error
select {
case sessionErr = <-errChan:
default:
}
assert.Equal(t, int32(2), dials.Load(), "expected no handover after the session ended")
if tt.wantErr == nil {
assert.NoError(t, sessionErr)
return
}
assert.ErrorIs(t, sessionErr, tt.wantErr)
assert.NotErrorIs(t, sessionErr, ErrHandoverFailed)
})
}
}

// TestClientExitsWhenServerCommandFails reproduces the missing-sshd case: the server accepts the
// websocket but can't launch its command, so it closes the connection immediately. The client
// proxy must exit promptly instead of hanging on the handover goroutine (which would leave the
Expand Down
Loading
Loading