[client] Deduplicate SSH client handshake and bound it with a deadline

Extract the dial-then-handshake sequence into nbssh.Handshake, which
applies the context deadline to the socket for the duration of the
handshake. Previously only the Android client did this; the CLI, wasm
and SSH proxy paths could block forever on a peer that accepts the TCP
connection and then goes silent, since ClientConfig.Timeout is not used
by NewClientConn.
This commit is contained in:
Zoltan Papp
2026-08-14 21:57:21 +02:00
parent 0bb49fa144
commit 1aa1f915a2
5 changed files with 61 additions and 37 deletions

View File

@@ -534,30 +534,11 @@ func (s *SSHClient) dialAndHandshake(gen uint64, host string, port int, clientCo
return fmt.Errorf("dial %s: %w", addr, err)
}
// DialContext bounds only the TCP establishment; without a deadline on the
// socket a peer that accepts and then goes silent blocks the handshake
// forever.
if deadline, ok := ctx.Deadline(); ok {
if err := conn.SetDeadline(deadline); err != nil {
closeQuiet(conn, "conn after deadline error")
return fmt.Errorf("set handshake deadline: %w", err)
}
}
sshConn, chans, reqs, err := gossh.NewClientConn(conn, addr, clientConfig)
client, err := nbssh.Handshake(ctx, conn, addr, clientConfig)
if err != nil {
if cerr := conn.Close(); cerr != nil {
log.Debugf("ssh: close after handshake error: %v", cerr)
}
return fmt.Errorf("ssh handshake: %w", err)
return err
}
if err := conn.SetDeadline(time.Time{}); err != nil {
closeQuiet(sshConn, "ssh conn after deadline clear error")
return fmt.Errorf("clear handshake deadline: %w", err)
}
client := gossh.NewClient(sshConn, chans, reqs)
s.mu.Lock()
if gen != s.gen {
s.mu.Unlock()

View File

@@ -313,21 +313,23 @@ func Dial(ctx context.Context, addr, user string, opts DialOptions) (*Client, er
// dialSSH establishes an SSH connection without JWT authentication
func dialSSH(ctx context.Context, network, addr string, config *ssh.ClientConfig) (*Client, error) {
if config.Timeout > 0 {
var cancel context.CancelFunc
ctx, cancel = context.WithTimeout(ctx, config.Timeout)
defer cancel()
}
dialer := &net.Dialer{}
conn, err := dialer.DialContext(ctx, network, addr)
if err != nil {
return nil, fmt.Errorf("dial %s: %w", addr, err)
}
clientConn, chans, reqs, err := ssh.NewClientConn(conn, addr, config)
client, err := nbssh.Handshake(ctx, conn, addr, config)
if err != nil {
if closeErr := conn.Close(); closeErr != nil {
log.Debugf("connection close after handshake failure: %v", closeErr)
}
return nil, fmt.Errorf("ssh handshake: %w", err)
return nil, err
}
client := ssh.NewClient(clientConn, chans, reqs)
return &Client{
client: client,
}, nil

45
client/ssh/handshake.go Normal file
View File

@@ -0,0 +1,45 @@
package ssh
import (
"context"
"fmt"
"io"
"net"
"time"
log "github.com/sirupsen/logrus"
"golang.org/x/crypto/ssh"
)
// Handshake runs the SSH client handshake on an already dialed conn and
// returns the resulting client. Dialing bounds only the TCP establishment;
// without a deadline on the socket a peer that accepts and then goes silent
// blocks the handshake forever, so the context deadline is applied to conn
// for the duration of the handshake. conn is closed on any error.
func Handshake(ctx context.Context, conn net.Conn, addr string, config *ssh.ClientConfig) (*ssh.Client, error) {
if deadline, ok := ctx.Deadline(); ok {
if err := conn.SetDeadline(deadline); err != nil {
closeHandshake(conn, "conn after deadline error")
return nil, fmt.Errorf("set handshake deadline: %w", err)
}
}
sshConn, chans, reqs, err := ssh.NewClientConn(conn, addr, config)
if err != nil {
closeHandshake(conn, "conn after handshake error")
return nil, fmt.Errorf("ssh handshake: %w", err)
}
if err := conn.SetDeadline(time.Time{}); err != nil {
closeHandshake(sshConn, "ssh conn after deadline clear error")
return nil, fmt.Errorf("clear handshake deadline: %w", err)
}
return ssh.NewClient(sshConn, chans, reqs), nil
}
func closeHandshake(c io.Closer, label string) {
if err := c.Close(); err != nil {
log.Debugf("ssh: close %s: %v", label, err)
}
}

View File

@@ -610,13 +610,10 @@ func (p *SSHProxy) dialBackend(ctx context.Context, addr, user, jwtToken string)
return nil, fmt.Errorf("connect to server: %w", err)
}
clientConn, chans, reqs, err := cryptossh.NewClientConn(conn, addr, config)
if err != nil {
_ = conn.Close()
return nil, fmt.Errorf("SSH handshake: %w", err)
}
handshakeCtx, cancel := context.WithTimeout(ctx, sshHandshakeTimeout)
defer cancel()
return cryptossh.NewClient(clientConn, chans, reqs), nil
return nbssh.Handshake(handshakeCtx, conn, addr, config)
}
func (p *SSHProxy) verifyHostKey(hostname string, remote net.Addr, key cryptossh.PublicKey) error {

View File

@@ -80,13 +80,12 @@ func (c *Client) Connect(host string, port int, username, jwtToken string, ipVer
return fmt.Errorf("dial %s: %w", addr, err)
}
sshConn, chans, reqs, err := ssh.NewClientConn(conn, addr, config)
sshClient, err := nbssh.Handshake(ctx, conn, addr, config)
if err != nil {
closeWithLog(conn, "connection after handshake error")
return fmt.Errorf("SSH handshake: %w", err)
return err
}
c.sshClient = ssh.NewClient(sshConn, chans, reqs)
c.sshClient = sshClient
logrus.Infof("SSH: Connected to %s", addr)
return nil