diff --git a/experimental/ssh/internal/proxy/client.go b/experimental/ssh/internal/proxy/client.go index 8d6a12fb438..b31c928865c 100644 --- a/experimental/ssh/internal/proxy/client.go +++ b/experimental/ssh/internal/proxy/client.go @@ -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. diff --git a/experimental/ssh/internal/proxy/client_server_test.go b/experimental/ssh/internal/proxy/client_server_test.go index 11a051e74c6..d215612c151 100644 --- a/experimental/ssh/internal/proxy/client_server_test.go +++ b/experimental/ssh/internal/proxy/client_server_test.go @@ -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 { @@ -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 { @@ -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 }, } } @@ -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 diff --git a/experimental/ssh/internal/proxy/handover_teardown_test.go b/experimental/ssh/internal/proxy/handover_teardown_test.go new file mode 100644 index 00000000000..46580d86b73 --- /dev/null +++ b/experimental/ssh/internal/proxy/handover_teardown_test.go @@ -0,0 +1,236 @@ +package proxy + +import ( + "context" + "io" + "net/http" + "net/http/httptest" + "sync/atomic" + "testing" + "time" + + "github.com/gorilla/websocket" + "github.com/stretchr/testify/require" +) + +// closeSignalingSource reports when the proxy closes it, which happens during the proxy's teardown. +type closeSignalingSource struct { + io.ReadCloser + closed chan struct{} +} + +func (s *closeSignalingSource) Close() error { + err := s.ReadCloser.Close() + close(s.closed) + return err +} + +// The receiving loop also runs on the server. If sshd's stdout ends while the server reads the +// client's normal close acknowledgment for a handover, the server must still finish the swap, so its +// teardown closes the replacement connection gracefully and the client sees a clean exit. +func TestServerEOFDuringHandoverIsACleanExit(t *testing.T) { + ctx, cancel := context.WithTimeout(t.Context(), 5*time.Second) + defer cancel() + serverProxy := newProxyConnection(nil) + serverInput, serverWriter := io.Pipe() + defer serverWriter.Close() + source := &closeSignalingSource{ReadCloser: serverInput, closed: make(chan struct{})} + serverReady := make(chan struct{}) + ackReceived := make(chan struct{}) + releaseAck := make(chan struct{}) + serverDone := make(chan error, 1) + handoverDone := make(chan error, 1) + var accepted atomic.Bool + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if !accepted.CompareAndSwap(false, true) { + handoverDone <- serverProxy.acceptHandover(ctx, w, r) + return + } + if err := serverProxy.accept(w, r); err != nil { + serverDone <- err + return + } + // Hold the receiving loop right after it reads the client's close acknowledgment, until + // the server's source has ended and its teardown has started. + conn := serverProxy.conn.Load() + defaultCloseHandler := conn.CloseHandler() + conn.SetCloseHandler(func(code int, text string) error { + err := defaultCloseHandler(code, text) + close(ackReceived) + select { + case <-releaseAck: + case <-ctx.Done(): + } + return err + }) + close(serverReady) + // Mirrors runServerProxy. + err := serverProxy.start(ctx, source, io.Discard) + closeProxyConnection(ctx, serverProxy) + serverDone <- err + })) + defer server.Close() + + clientInput, clientWriter := io.Pipe() + defer clientWriter.Close() + clientOutput := newTestBuffer(t) + ticks := make(chan time.Time, 1) + clientDone := make(chan error, 1) + go func() { + clientDone <- RunClientProxy(ctx, clientInput, clientOutput, func() <-chan time.Time { return ticks }, time.Hour, false, + func(dialCtx context.Context, _ DialRequest) (*websocket.Conn, error) { + conn, _, err := websocket.DefaultDialer.DialContext(dialCtx, "ws"+server.URL[4:], nil) // nolint:bodyclose + return conn, err + }) + }() + + select { + case <-serverReady: + case <-ctx.Done(): + t.Fatal("the server did not accept the connection") + } + // The client waits for the server's first byte before it treats the session as started. Wait + // until the client has it: the write only proves the sending loop read it, and a handover that + // takes the mutex before the sending loop sends it would block the loop from reading the EOF. + banner := []byte("SSH-2.0-test\r\n") + _, err := serverWriter.Write(banner) + require.NoError(t, err) + require.NoError(t, clientOutput.WaitForWrite(banner)) + ticks <- time.Now() + select { + case <-ackReceived: + case <-ctx.Done(): + t.Fatal("the server did not receive the handover close acknowledgment") + } + require.NoError(t, serverWriter.Close()) + select { + case <-source.closed: + case <-ctx.Done(): + t.Fatal("the server did not start its teardown") + } + close(releaseAck) + + select { + case <-handoverDone: + case <-ctx.Done(): + t.Fatal("the server handover did not finish") + } + select { + case err := <-serverDone: + require.NoError(t, err) + case <-ctx.Done(): + t.Fatal("the server did not finish") + } + select { + case err := <-clientDone: + require.NoError(t, err, "a clean server exit during a handover must be a clean client exit") + case <-ctx.Done(): + t.Fatal("the client did not finish") + } +} + +// If the client's input ends while its receiving loop holds the server's normal handover close, the +// client must still complete the swap: its exit then sends the resumable "finished" close on the +// replacement connection, and the server ends the session at once instead of waiting for a reattach. +func TestClientEOFDuringHandoverReleasesServer(t *testing.T) { + ctx, cancel := context.WithTimeout(t.Context(), 5*time.Second) + defer cancel() + serverProxy := newResumableProxyConnection(nil) + serverInput, serverWriter := io.Pipe() + defer serverWriter.Close() + serverReady := make(chan struct{}) + serverDone := make(chan error, 1) + handoverDone := make(chan error, 1) + var accepted atomic.Bool + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if !accepted.CompareAndSwap(false, true) { + handoverDone <- serverProxy.acceptHandover(ctx, w, r) + return + } + if err := serverProxy.accept(w, r); err != nil { + serverDone <- err + return + } + close(serverReady) + // Mirrors runServerProxy. + err := serverProxy.start(ctx, serverInput, io.Discard) + closeProxyConnection(ctx, serverProxy) + serverDone <- err + })) + defer server.Close() + + clientInput, clientWriter := io.Pipe() + defer clientWriter.Close() + clientSource := &closeSignalingSource{ReadCloser: clientInput, closed: make(chan struct{})} + clientOutput := newTestBuffer(t) + ticks := make(chan time.Time, 1) + clientDone := make(chan error, 1) + ackSent := make(chan struct{}) + releaseAck := make(chan struct{}) + var dials atomic.Int32 + go func() { + clientDone <- RunClientProxy(ctx, clientSource, clientOutput, func() <-chan time.Time { return ticks }, time.Hour, true, + func(dialCtx context.Context, _ DialRequest) (*websocket.Conn, error) { + conn, _, err := websocket.DefaultDialer.DialContext(dialCtx, "ws"+server.URL[4:], nil) // nolint:bodyclose + // Hold the client's receiving loop right after it acknowledges the server's + // handover close on the first connection, until the client's teardown has started. + if err == nil && dials.Add(1) == 1 { + defaultCloseHandler := conn.CloseHandler() + conn.SetCloseHandler(func(code int, text string) error { + err := defaultCloseHandler(code, text) + close(ackSent) + select { + case <-releaseAck: + case <-ctx.Done(): + } + return err + }) + } + return conn, err + }) + }() + + select { + case <-serverReady: + case <-ctx.Done(): + t.Fatal("the server did not accept the connection") + } + banner := []byte("SSH-2.0-test\r\n") + _, err := serverWriter.Write(banner) + require.NoError(t, err) + require.NoError(t, clientOutput.WaitForWrite(banner)) + ticks <- time.Now() + select { + case <-ackSent: + case <-ctx.Done(): + t.Fatal("the client did not acknowledge the handover close") + } + select { + case err := <-handoverDone: + require.NoError(t, err) + case <-ctx.Done(): + t.Fatal("the server did not finish the handover") + } + require.NoError(t, clientWriter.Close()) + select { + case <-clientSource.closed: + case <-ctx.Done(): + t.Fatal("the client did not start its teardown") + } + close(releaseAck) + + select { + case err := <-clientDone: + require.NoError(t, err) + case <-ctx.Done(): + t.Fatal("the client did not finish") + } + select { + case err := <-serverDone: + require.NoError(t, err) + case <-serverProxy.resume.parked: + t.Fatalf("a clean client exit left the server waiting %v for a reattach", serverProxy.resume.grace) + case <-ctx.Done(): + t.Fatal("the server did not finish after a clean client exit") + } +} diff --git a/experimental/ssh/internal/proxy/proxy.go b/experimental/ssh/internal/proxy/proxy.go index 493644ad5ae..88cee6eff3a 100644 --- a/experimental/ssh/internal/proxy/proxy.go +++ b/experimental/ssh/internal/proxy/proxy.go @@ -535,7 +535,18 @@ func (pc *proxyConnection) runReceivingLoop(ctx context.Context, dst io.Writer) // reattachment to recover bytes the close-frame exchange did not drain. if handover := pc.handoverState.Load(); handover != nil { var closeConnSignal error - if !websocket.IsCloseError(err, websocket.CloseNormalClosure) { + switch { + case websocket.IsCloseError(err, websocket.CloseNormalClosure): + // A normal close ends the old connection cleanly, usually as the peer's side of + // the handover. Complete the swap even when the session is ending: teardown then + // closes the replacement connection gracefully, and a resumable peer does not + // wait for a reattach. + case ctx.Err() != nil: + // The session is ending and teardown closed the connection under the handover + // (see the ctx.Err() check below). Pass the cancellation on, so the handover + // initiator does not report a failed handover for an ordinary exit. + closeConnSignal = ctx.Err() + default: closeConnSignal = errors.Join(ErrWebsocketDropped, fmt.Errorf("failed to read from websocket during handover: %w", err)) } // Signal the current connection is closed to the handover initiator (initiateHandover or acceptHandover).