[android] invalidate stale SSH operations across reconnects

This commit is contained in:
Zoltán Papp
2026-08-12 12:52:24 +02:00
parent ba16475ad6
commit 9531c9cf79
+53 -8
View File
@@ -44,7 +44,10 @@ const PasswordRequiredMarker = "netbird-ssh-password-required"
// reach this: NetBird peers verify against the registry. // reach this: NetBird peers verify against the registry.
const HostKeyUnknownMarker = "netbird-ssh-hostkey-unknown" const HostKeyUnknownMarker = "netbird-ssh-hostkey-unknown"
var errPasswordRequired = errors.New(PasswordRequiredMarker) var (
errPasswordRequired = errors.New(PasswordRequiredMarker)
errClientClosed = errors.New("ssh client closed")
)
// errHostKeyUnknown carries the presented fingerprint so Connect can build the // errHostKeyUnknown carries the presented fingerprint so Connect can build the
// marker message the Java side parses. // marker message the Java side parses.
@@ -96,6 +99,13 @@ type SSHClient struct {
stdin io.WriteCloser stdin io.WriteCloser
closed bool closed bool
// gen identifies the current connection attempt. Connect and Close bump it,
// so an in-flight dial or a reader left over from a previous connection
// finds itself stale and stays silent instead of publishing OnConnected or
// OnClose for a connection the caller already abandoned.
gen uint64
dialCancel context.CancelFunc
// knownHostsPath is the TOFU store for regular SSH servers. Java supplies a // knownHostsPath is the TOFU store for regular SSH servers. Java supplies a
// per-profile path, since an overlay IP is a different host under a // per-profile path, since an overlay IP is a different host under a
// different profile. Empty until set: without it a regular server cannot be // different profile. Empty until set: without it a regular server cannot be
@@ -180,6 +190,11 @@ func (s *SSHClient) Connect(host string, port int, user, password string) error
return errors.New("netbird engine not available") return errors.New("netbird engine not available")
} }
s.mu.Lock()
s.gen++
gen := s.gen
s.mu.Unlock()
serverType := detectServerType(host, port) serverType := detectServerType(host, port)
log.Debugf("SSH server type: %s", serverType) log.Debugf("SSH server type: %s", serverType)
@@ -194,7 +209,7 @@ func (s *SSHClient) Connect(host string, port int, user, password string) error
HostKeyCallback: hostKeyCallback, HostKeyCallback: hostKeyCallback,
Timeout: sshDialTimeout, Timeout: sshDialTimeout,
} }
err = s.dialAndHandshake(host, port, clientConfig) err = s.dialAndHandshake(gen, host, port, clientConfig)
// An unknown host key is a prompt, not a failure: return the marker intact // An unknown host key is a prompt, not a failure: return the marker intact
// (rootCause would unwrap it) so Java can show the fingerprint and retry. // (rootCause would unwrap it) so Java can show the fingerprint and retry.
@@ -257,6 +272,7 @@ func (s *SSHClient) startSession(cols, rows int) error {
log.Debugf("SSH: starting session %dx%d", cols, rows) log.Debugf("SSH: starting session %dx%d", cols, rows)
s.mu.Lock() s.mu.Lock()
sshClient := s.sshClient sshClient := s.sshClient
gen := s.gen
s.mu.Unlock() s.mu.Unlock()
if sshClient == nil { if sshClient == nil {
@@ -303,6 +319,11 @@ func (s *SSHClient) startSession(cols, rows int) error {
} }
s.mu.Lock() s.mu.Lock()
if gen != s.gen {
s.mu.Unlock()
closeQuiet(session, "stale session")
return errClientClosed
}
s.session = session s.session = session
s.stdin = stdin s.stdin = stdin
s.mu.Unlock() s.mu.Unlock()
@@ -315,7 +336,7 @@ func (s *SSHClient) startSession(cols, rows int) error {
if second := <-readerDone; reason == "" { if second := <-readerDone; reason == "" {
reason = second reason = second
} }
s.notifyClose(reason) s.notifyClose(gen, reason)
}() }()
log.Debug("SSH: session started, shell running") log.Debug("SSH: session started, shell running")
return nil return nil
@@ -358,12 +379,20 @@ func (s *SSHClient) Reset() {
// multiple times. // multiple times.
func (s *SSHClient) Close() error { func (s *SSHClient) Close() error {
s.mu.Lock() s.mu.Lock()
s.gen++
if s.dialCancel != nil {
s.dialCancel()
s.dialCancel = nil
}
sshClient := s.sshClient sshClient := s.sshClient
session := s.session session := s.session
stdin := s.stdin stdin := s.stdin
s.sshClient = nil s.sshClient = nil
s.session = nil s.session = nil
s.stdin = nil s.stdin = nil
notify := !s.closed
s.closed = true
listener := s.listener
s.mu.Unlock() s.mu.Unlock()
if stdin != nil { if stdin != nil {
@@ -382,7 +411,9 @@ func (s *SSHClient) Close() error {
firstErr = err firstErr = err
} }
} }
s.notifyClose("closed by client") if notify && listener != nil {
listener.OnClose("closed by client")
}
return firstErr return firstErr
} }
@@ -556,11 +587,19 @@ func (s *SSHClient) requestJWTToken(cfg *profilemanager.Config) (string, error)
return token, nil return token, nil
} }
func (s *SSHClient) dialAndHandshake(host string, port int, clientConfig *gossh.ClientConfig) error { func (s *SSHClient) dialAndHandshake(gen uint64, host string, port int, clientConfig *gossh.ClientConfig) error {
addr := net.JoinHostPort(host, strconv.Itoa(port)) addr := net.JoinHostPort(host, strconv.Itoa(port))
ctx, cancel := context.WithTimeout(context.Background(), sshDialTimeout) ctx, cancel := context.WithTimeout(context.Background(), sshDialTimeout)
defer cancel() defer cancel()
s.mu.Lock()
if gen != s.gen {
s.mu.Unlock()
return errClientClosed
}
s.dialCancel = cancel
s.mu.Unlock()
var dialer net.Dialer var dialer net.Dialer
conn, err := dialer.DialContext(ctx, "tcp", addr) conn, err := dialer.DialContext(ctx, "tcp", addr)
if err != nil { if err != nil {
@@ -590,8 +629,14 @@ func (s *SSHClient) dialAndHandshake(host string, port int, clientConfig *gossh.
return fmt.Errorf("clear handshake deadline: %w", err) return fmt.Errorf("clear handshake deadline: %w", err)
} }
client := gossh.NewClient(sshConn, chans, reqs)
s.mu.Lock() s.mu.Lock()
s.sshClient = gossh.NewClient(sshConn, chans, reqs) if gen != s.gen {
s.mu.Unlock()
closeQuiet(client, "stale ssh client")
return errClientClosed
}
s.sshClient = client
listener := s.listener listener := s.listener
s.mu.Unlock() s.mu.Unlock()
@@ -637,9 +682,9 @@ func (s *SSHClient) notifyStatus(text string) {
} }
} }
func (s *SSHClient) notifyClose(reason string) { func (s *SSHClient) notifyClose(gen uint64, reason string) {
s.mu.Lock() s.mu.Lock()
if s.closed { if gen != s.gen || s.closed {
s.mu.Unlock() s.mu.Unlock()
return return
} }