From 2da4512272f16b9573ff4717feaaeeb463b370ae Mon Sep 17 00:00:00 2001 From: Zoltan Papp Date: Fri, 14 Aug 2026 17:49:31 +0200 Subject: [PATCH] [client] Deduplicate SSH PTY session setup and host key verification Extract the identical PTY session setup shared by the wasm and Android terminal clients into ssh.StartPTYSession, and move the stored-key host verification onto the engine so the embed client delegates and the Android client passes the engine directly as HostKeyVerifier. Co-Authored-By: Claude Fable 5 --- client/android/ssh_client.go | 65 ++++---------------------- client/embed/embed.go | 8 +--- client/internal/engine_ssh.go | 11 +++++ client/ssh/session.go | 73 ++++++++++++++++++++++++++++++ client/wasm/internal/ssh/client.go | 46 +++---------------- 5 files changed, 100 insertions(+), 103 deletions(-) create mode 100644 client/ssh/session.go diff --git a/client/android/ssh_client.go b/client/android/ssh_client.go index e76250280..1c4015906 100644 --- a/client/android/ssh_client.go +++ b/client/android/ssh_client.go @@ -59,19 +59,6 @@ func (e *errHostKeyUnknown) Error() string { return HostKeyUnknownMarker + ":" + e.fingerprint } -// engineHostKeyVerifier adapts *internal.Engine to nbssh.HostKeyVerifier. -type engineHostKeyVerifier struct { - engine *internal.Engine -} - -func (v *engineHostKeyVerifier) VerifySSHHostKey(peerAddress string, presented []byte) error { - storedKey, found := v.engine.GetPeerSSHKey(peerAddress) - if !found { - return nbssh.ErrPeerNotFound - } - return nbssh.VerifyHostKey(storedKey, presented, peerAddress) -} - // SSHTerminalListener receives SSH session events. It is implemented in Java. // // All callbacks are invoked from goroutines and may run concurrently with each @@ -329,58 +316,24 @@ func (s *SSHClient) startSession(cols, rows int) error { return errors.New("ssh client not connected") } - session, err := sshClient.NewSession() + pty, err := nbssh.StartPTYSession(sshClient, cols, rows) if err != nil { - return fmt.Errorf("new session: %w", err) - } - - modes := gossh.TerminalModes{ - gossh.ECHO: 1, - gossh.TTY_OP_ISPEED: 14400, - gossh.TTY_OP_OSPEED: 14400, - gossh.VINTR: 3, - gossh.VQUIT: 28, - gossh.VERASE: 127, - } - if err := session.RequestPty("xterm-256color", rows, cols, modes); err != nil { - closeQuiet(session, "session after pty error") - return fmt.Errorf("request pty: %w", err) - } - - stdin, err := session.StdinPipe() - if err != nil { - closeQuiet(session, "session after stdin error") - return fmt.Errorf("stdin pipe: %w", err) - } - stdout, err := session.StdoutPipe() - if err != nil { - closeQuiet(session, "session after stdout error") - return fmt.Errorf("stdout pipe: %w", err) - } - stderr, err := session.StderrPipe() - if err != nil { - closeQuiet(session, "session after stderr error") - return fmt.Errorf("stderr pipe: %w", err) - } - - if err := session.Shell(); err != nil { - closeQuiet(session, "session after shell error") - return fmt.Errorf("start shell: %w", err) + return err } s.mu.Lock() if gen != s.gen { s.mu.Unlock() - closeQuiet(session, "stale session") + closeQuiet(pty.Session, "stale session") return errClientClosed } - s.session = session - s.stdin = stdin + s.session = pty.Session + s.stdin = pty.Stdin s.mu.Unlock() readerDone := make(chan string, 2) - go func() { readerDone <- s.readLoop(stdout, "stdout") }() - go func() { readerDone <- s.readLoop(stderr, "stderr") }() + go func() { readerDone <- s.readLoop(pty.Stdout, "stdout") }() + go func() { readerDone <- s.readLoop(pty.Stderr, "stderr") }() go func() { reason := <-readerDone if second := <-readerDone; reason == "" { @@ -402,7 +355,7 @@ func (s *SSHClient) buildAuth(cfg *profilemanager.Config, engine *internal.Engin return nil, nil, fmt.Errorf("jwt: %w", err) } auths := []gossh.AuthMethod{gossh.Password(token)} - return auths, nbssh.CreateHostKeyCallback(&engineHostKeyVerifier{engine: engine}), nil + return auths, nbssh.CreateHostKeyCallback(engine), nil case detection.ServerTypeNetBirdNoJWT: if cfg.SSHKey == "" { @@ -413,7 +366,7 @@ func (s *SSHClient) buildAuth(cfg *profilemanager.Config, engine *internal.Engin return nil, nil, fmt.Errorf("parse netbird ssh key: %w", err) } auths := []gossh.AuthMethod{gossh.PublicKeys(signer)} - return auths, nbssh.CreateHostKeyCallback(&engineHostKeyVerifier{engine: engine}), nil + return auths, nbssh.CreateHostKeyCallback(engine), nil case detection.ServerTypeRegular: var auths []gossh.AuthMethod diff --git a/client/embed/embed.go b/client/embed/embed.go index 99a6b8229..6a3c25c33 100644 --- a/client/embed/embed.go +++ b/client/embed/embed.go @@ -21,7 +21,6 @@ import ( "github.com/netbirdio/netbird/client/internal/auth" "github.com/netbirdio/netbird/client/internal/peer" "github.com/netbirdio/netbird/client/internal/profilemanager" - sshcommon "github.com/netbirdio/netbird/client/ssh" "github.com/netbirdio/netbird/client/system" "github.com/netbirdio/netbird/shared/management/domain" mgmProto "github.com/netbirdio/netbird/shared/management/proto" @@ -521,12 +520,7 @@ func (c *Client) VerifySSHHostKey(peerAddress string, key []byte) error { return err } - storedKey, found := engine.GetPeerSSHKey(peerAddress) - if !found { - return sshcommon.ErrPeerNotFound - } - - return sshcommon.VerifyHostKey(storedKey, key, peerAddress) + return engine.VerifySSHHostKey(peerAddress, key) } // SetPerformance retunes a running Client. Only PreallocatedBuffersPerPool diff --git a/client/internal/engine_ssh.go b/client/internal/engine_ssh.go index 53d2c1122..5c86884db 100644 --- a/client/internal/engine_ssh.go +++ b/client/internal/engine_ssh.go @@ -12,6 +12,7 @@ import ( firewallManager "github.com/netbirdio/netbird/client/firewall/manager" "github.com/netbirdio/netbird/client/iface/netstack" nftypes "github.com/netbirdio/netbird/client/internal/netflow/types" + nbssh "github.com/netbirdio/netbird/client/ssh" sshauth "github.com/netbirdio/netbird/client/ssh/auth" sshconfig "github.com/netbirdio/netbird/client/ssh/config" sshserver "github.com/netbirdio/netbird/client/ssh/server" @@ -216,6 +217,16 @@ func (e *Engine) GetPeerSSHKey(peerAddress string) ([]byte, bool) { return nil, false } +// VerifySSHHostKey verifies a presented SSH host key against the stored key of +// the peer at peerAddress. It implements ssh.HostKeyVerifier. +func (e *Engine) VerifySSHHostKey(peerAddress string, presentedKey []byte) error { + storedKey, found := e.GetPeerSSHKey(peerAddress) + if !found { + return nbssh.ErrPeerNotFound + } + return nbssh.VerifyHostKey(storedKey, presentedKey, peerAddress) +} + // cleanupSSHConfig removes NetBird SSH client configuration on shutdown func (e *Engine) cleanupSSHConfig() { if netstack.IsEnabled() { diff --git a/client/ssh/session.go b/client/ssh/session.go new file mode 100644 index 000000000..f0faea023 --- /dev/null +++ b/client/ssh/session.go @@ -0,0 +1,73 @@ +package ssh + +import ( + "fmt" + "io" + + log "github.com/sirupsen/logrus" + "golang.org/x/crypto/ssh" +) + +// defaultTerminalModes are the PTY modes used by the interactive terminal clients. +var defaultTerminalModes = ssh.TerminalModes{ + ssh.ECHO: 1, + ssh.TTY_OP_ISPEED: 14400, + ssh.TTY_OP_OSPEED: 14400, + ssh.VINTR: 3, + ssh.VQUIT: 28, + ssh.VERASE: 127, +} + +// PTYSession is an interactive shell session with a PTY and its I/O pipes. +type PTYSession struct { + Session *ssh.Session + Stdin io.WriteCloser + Stdout io.Reader + Stderr io.Reader +} + +// StartPTYSession opens a session on the client, requests an xterm-256color PTY +// with the default terminal modes, wires up the I/O pipes and starts a shell. +// The session is closed on any error. +func StartPTYSession(client *ssh.Client, cols, rows int) (*PTYSession, error) { + session, err := client.NewSession() + if err != nil { + return nil, fmt.Errorf("new session: %w", err) + } + + pty, err := setupPTYSession(session, cols, rows) + if err != nil { + if closeErr := session.Close(); closeErr != nil { + log.Debugf("ssh: session close after setup error: %v", closeErr) + } + return nil, err + } + return pty, nil +} + +// setupPTYSession requests the PTY, opens the pipes and starts the shell on an +// already created session. +func setupPTYSession(session *ssh.Session, cols, rows int) (*PTYSession, error) { + if err := session.RequestPty("xterm-256color", rows, cols, defaultTerminalModes); err != nil { + return nil, fmt.Errorf("request pty: %w", err) + } + + stdin, err := session.StdinPipe() + if err != nil { + return nil, fmt.Errorf("stdin pipe: %w", err) + } + stdout, err := session.StdoutPipe() + if err != nil { + return nil, fmt.Errorf("stdout pipe: %w", err) + } + stderr, err := session.StderrPipe() + if err != nil { + return nil, fmt.Errorf("stderr pipe: %w", err) + } + + if err := session.Shell(); err != nil { + return nil, fmt.Errorf("start shell: %w", err) + } + + return &PTYSession{Session: session, Stdin: stdin, Stdout: stdout, Stderr: stderr}, nil +} diff --git a/client/wasm/internal/ssh/client.go b/client/wasm/internal/ssh/client.go index 9cfe65266..83170bfc3 100644 --- a/client/wasm/internal/ssh/client.go +++ b/client/wasm/internal/ssh/client.go @@ -125,51 +125,17 @@ func (c *Client) StartSession(cols, rows int) error { return fmt.Errorf("SSH client not connected") } - session, err := c.sshClient.NewSession() + pty, err := nbssh.StartPTYSession(c.sshClient, cols, rows) if err != nil { - return fmt.Errorf("create session: %w", err) + return err } c.mu.Lock() defer c.mu.Unlock() - c.session = session - - modes := ssh.TerminalModes{ - ssh.ECHO: 1, - ssh.TTY_OP_ISPEED: 14400, - ssh.TTY_OP_OSPEED: 14400, - ssh.VINTR: 3, - ssh.VQUIT: 28, - ssh.VERASE: 127, - } - - if err := session.RequestPty("xterm-256color", rows, cols, modes); err != nil { - closeWithLog(session, "session after PTY error") - return fmt.Errorf("PTY request: %w", err) - } - - c.stdin, err = session.StdinPipe() - if err != nil { - closeWithLog(session, "session after stdin error") - return fmt.Errorf("get stdin: %w", err) - } - - c.stdout, err = session.StdoutPipe() - if err != nil { - closeWithLog(session, "session after stdout error") - return fmt.Errorf("get stdout: %w", err) - } - - c.stderr, err = session.StderrPipe() - if err != nil { - closeWithLog(session, "session after stderr error") - return fmt.Errorf("get stderr: %w", err) - } - - if err := session.Shell(); err != nil { - closeWithLog(session, "session after shell error") - return fmt.Errorf("start shell: %w", err) - } + c.session = pty.Session + c.stdin = pty.Stdin + c.stdout = pty.Stdout + c.stderr = pty.Stderr logrus.Info("SSH: Session started with PTY") return nil