mirror of
https://github.com/netbirdio/netbird.git
synced 2026-10-09 15:09:08 +02:00
[android] invalidate stale SSH operations across reconnects
This commit is contained in:
@@ -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
|
||||||
}
|
}
|
||||||
|
|||||||
Reference in New Issue
Block a user