diff --git a/pkg/ssh/forward.go b/pkg/ssh/forward.go index f273e20e2..f87d3ce10 100644 --- a/pkg/ssh/forward.go +++ b/pkg/ssh/forward.go @@ -18,6 +18,8 @@ import ( // should check for this error with errors.Is. var ErrIdleTimeout = errors.New("port forward idle timeout") +var ErrTransportClosed = errors.New("ssh transport closed") + type ForwardingFunction func( net.Conn, *ssh.Client, @@ -114,6 +116,23 @@ func portForwarding( } }() + if client != nil { + transportClosed := make(chan struct{}) + go func() { + if werr := client.Wait(); werr != nil { + log.Debugf("ssh transport closed on %s: %v", srcAddr, werr) + } + close(transportClosed) + }() + go func() { + select { + case <-done: + case <-transportClosed: + cancel(ErrTransportClosed) + } + }() + } + counter := newConnectionCounter(fwdCtx, exitAfterTimeout, func() { log.Infof( "Stopping port-forward on %s: idle for a while. "+ @@ -126,10 +145,11 @@ func portForwarding( // waiting for a new connection connection, err := listener.Accept() if err != nil { - // If shutdown was caused by the idle timeout, surface that - // typed error so callers can choose to treat it as a clean exit. - if cause := context.Cause(fwdCtx); errors.Is(cause, ErrIdleTimeout) { + switch cause := context.Cause(fwdCtx); { + case errors.Is(cause, ErrIdleTimeout): return ErrIdleTimeout + case errors.Is(cause, ErrTransportClosed): + return ErrTransportClosed } return err } diff --git a/pkg/ssh/forward_test.go b/pkg/ssh/forward_test.go index 6411f36d7..78152de83 100644 --- a/pkg/ssh/forward_test.go +++ b/pkg/ssh/forward_test.go @@ -2,6 +2,7 @@ package ssh import ( "context" + "crypto/ed25519" "errors" "net" "testing" @@ -91,3 +92,95 @@ func TestPortForwarding_ParentCancelNotIdleTimeout(t *testing.T) { t.Fatal("timed out waiting for portForwarding to return") } } + +func startTestSSHServer(t *testing.T) (client *ssh.Client, kill func()) { + t.Helper() + + _, priv, err := ed25519.GenerateKey(nil) + if err != nil { + t.Fatalf("generate host key: %v", err) + } + signer, err := ssh.NewSignerFromKey(priv) + if err != nil { + t.Fatalf("signer: %v", err) + } + + srvCfg := &ssh.ServerConfig{NoClientAuth: true} + srvCfg.AddHostKey(signer) + + ln, err := net.Listen("tcp", "127.0.0.1:0") + if err != nil { + t.Fatalf("listen: %v", err) + } + t.Cleanup(func() { _ = ln.Close() }) + + serverConnCh := make(chan net.Conn, 1) + go func() { + nConn, aerr := ln.Accept() + if aerr != nil { + return + } + serverConnCh <- nConn + sConn, chans, reqs, herr := ssh.NewServerConn(nConn, srvCfg) + if herr != nil { + return + } + go ssh.DiscardRequests(reqs) + go func() { + for nc := range chans { + _ = nc.Reject(ssh.Prohibited, "no channels in test") + } + }() + _ = sConn + }() + + cConn, err := ssh.Dial("tcp", ln.Addr().String(), &ssh.ClientConfig{ + User: "test", + HostKeyCallback: ssh.FixedHostKey(signer.PublicKey()), + Timeout: 2 * time.Second, + }) + if err != nil { + t.Fatalf("dial: %v", err) + } + + serverConn := <-serverConnCh + kill = func() { + _ = serverConn.Close() + _ = ln.Close() + } + t.Cleanup(func() { _ = cConn.Close(); kill() }) + return cConn, kill +} + +func TestPortForwarding_ReleasesListenerOnTransportDeath(t *testing.T) { + client, kill := startTestSSHServer(t) + + lis, err := net.Listen("tcp", "127.0.0.1:0") + if err != nil { + t.Fatalf("listen: %v", err) + } + addr := lis.Addr().String() + + done := make(chan error, 1) + go func() { + done <- portForwarding(t.Context(), client, lis, addr, "tcp", "127.0.0.1:9", 0, forward) + }() + + time.Sleep(100 * time.Millisecond) + kill() + + select { + case err := <-done: + if !errors.Is(err, ErrTransportClosed) { + t.Fatalf("expected ErrTransportClosed, got %v", err) + } + case <-time.After(3 * time.Second): + t.Fatal("portForwarding did not return after transport death (listener leaked)") + } + + l2, err := net.Listen("tcp", addr) + if err != nil { + t.Fatalf("rebind on %s failed after transport death: %v", addr, err) + } + _ = l2.Close() +}