Skip to content
Merged
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
26 changes: 23 additions & 3 deletions pkg/ssh/forward.go
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -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. "+
Expand All @@ -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
}
Expand Down
93 changes: 93 additions & 0 deletions pkg/ssh/forward_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -2,6 +2,7 @@ package ssh

import (
"context"
"crypto/ed25519"
"errors"
"net"
"testing"
Expand Down Expand Up @@ -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)
}
Comment thread
coderabbitai[bot] marked this conversation as resolved.

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()
}
Loading