[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 <noreply@anthropic.com>
This commit is contained in:
Zoltan Papp
2026-08-14 17:49:31 +02:00
co-authored by Claude Fable 5
parent 16f7e1e148
commit 2da4512272
5 changed files with 100 additions and 103 deletions
+9 -56
View File
@@ -59,19 +59,6 @@ func (e *errHostKeyUnknown) Error() string {
return HostKeyUnknownMarker + ":" + e.fingerprint 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. // SSHTerminalListener receives SSH session events. It is implemented in Java.
// //
// All callbacks are invoked from goroutines and may run concurrently with each // 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") return errors.New("ssh client not connected")
} }
session, err := sshClient.NewSession() pty, err := nbssh.StartPTYSession(sshClient, cols, rows)
if err != nil { if err != nil {
return fmt.Errorf("new session: %w", err) return 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)
} }
s.mu.Lock() s.mu.Lock()
if gen != s.gen { if gen != s.gen {
s.mu.Unlock() s.mu.Unlock()
closeQuiet(session, "stale session") closeQuiet(pty.Session, "stale session")
return errClientClosed return errClientClosed
} }
s.session = session s.session = pty.Session
s.stdin = stdin s.stdin = pty.Stdin
s.mu.Unlock() s.mu.Unlock()
readerDone := make(chan string, 2) readerDone := make(chan string, 2)
go func() { readerDone <- s.readLoop(stdout, "stdout") }() go func() { readerDone <- s.readLoop(pty.Stdout, "stdout") }()
go func() { readerDone <- s.readLoop(stderr, "stderr") }() go func() { readerDone <- s.readLoop(pty.Stderr, "stderr") }()
go func() { go func() {
reason := <-readerDone reason := <-readerDone
if second := <-readerDone; reason == "" { 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) return nil, nil, fmt.Errorf("jwt: %w", err)
} }
auths := []gossh.AuthMethod{gossh.Password(token)} auths := []gossh.AuthMethod{gossh.Password(token)}
return auths, nbssh.CreateHostKeyCallback(&engineHostKeyVerifier{engine: engine}), nil return auths, nbssh.CreateHostKeyCallback(engine), nil
case detection.ServerTypeNetBirdNoJWT: case detection.ServerTypeNetBirdNoJWT:
if cfg.SSHKey == "" { 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) return nil, nil, fmt.Errorf("parse netbird ssh key: %w", err)
} }
auths := []gossh.AuthMethod{gossh.PublicKeys(signer)} auths := []gossh.AuthMethod{gossh.PublicKeys(signer)}
return auths, nbssh.CreateHostKeyCallback(&engineHostKeyVerifier{engine: engine}), nil return auths, nbssh.CreateHostKeyCallback(engine), nil
case detection.ServerTypeRegular: case detection.ServerTypeRegular:
var auths []gossh.AuthMethod var auths []gossh.AuthMethod
+1 -7
View File
@@ -21,7 +21,6 @@ import (
"github.com/netbirdio/netbird/client/internal/auth" "github.com/netbirdio/netbird/client/internal/auth"
"github.com/netbirdio/netbird/client/internal/peer" "github.com/netbirdio/netbird/client/internal/peer"
"github.com/netbirdio/netbird/client/internal/profilemanager" "github.com/netbirdio/netbird/client/internal/profilemanager"
sshcommon "github.com/netbirdio/netbird/client/ssh"
"github.com/netbirdio/netbird/client/system" "github.com/netbirdio/netbird/client/system"
"github.com/netbirdio/netbird/shared/management/domain" "github.com/netbirdio/netbird/shared/management/domain"
mgmProto "github.com/netbirdio/netbird/shared/management/proto" mgmProto "github.com/netbirdio/netbird/shared/management/proto"
@@ -521,12 +520,7 @@ func (c *Client) VerifySSHHostKey(peerAddress string, key []byte) error {
return err return err
} }
storedKey, found := engine.GetPeerSSHKey(peerAddress) return engine.VerifySSHHostKey(peerAddress, key)
if !found {
return sshcommon.ErrPeerNotFound
}
return sshcommon.VerifyHostKey(storedKey, key, peerAddress)
} }
// SetPerformance retunes a running Client. Only PreallocatedBuffersPerPool // SetPerformance retunes a running Client. Only PreallocatedBuffersPerPool
+11
View File
@@ -12,6 +12,7 @@ import (
firewallManager "github.com/netbirdio/netbird/client/firewall/manager" firewallManager "github.com/netbirdio/netbird/client/firewall/manager"
"github.com/netbirdio/netbird/client/iface/netstack" "github.com/netbirdio/netbird/client/iface/netstack"
nftypes "github.com/netbirdio/netbird/client/internal/netflow/types" nftypes "github.com/netbirdio/netbird/client/internal/netflow/types"
nbssh "github.com/netbirdio/netbird/client/ssh"
sshauth "github.com/netbirdio/netbird/client/ssh/auth" sshauth "github.com/netbirdio/netbird/client/ssh/auth"
sshconfig "github.com/netbirdio/netbird/client/ssh/config" sshconfig "github.com/netbirdio/netbird/client/ssh/config"
sshserver "github.com/netbirdio/netbird/client/ssh/server" sshserver "github.com/netbirdio/netbird/client/ssh/server"
@@ -216,6 +217,16 @@ func (e *Engine) GetPeerSSHKey(peerAddress string) ([]byte, bool) {
return nil, false 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 // cleanupSSHConfig removes NetBird SSH client configuration on shutdown
func (e *Engine) cleanupSSHConfig() { func (e *Engine) cleanupSSHConfig() {
if netstack.IsEnabled() { if netstack.IsEnabled() {
+73
View File
@@ -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
}
+6 -40
View File
@@ -125,51 +125,17 @@ func (c *Client) StartSession(cols, rows int) error {
return fmt.Errorf("SSH client not connected") return fmt.Errorf("SSH client not connected")
} }
session, err := c.sshClient.NewSession() pty, err := nbssh.StartPTYSession(c.sshClient, cols, rows)
if err != nil { if err != nil {
return fmt.Errorf("create session: %w", err) return err
} }
c.mu.Lock() c.mu.Lock()
defer c.mu.Unlock() defer c.mu.Unlock()
c.session = session c.session = pty.Session
c.stdin = pty.Stdin
modes := ssh.TerminalModes{ c.stdout = pty.Stdout
ssh.ECHO: 1, c.stderr = pty.Stderr
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)
}
logrus.Info("SSH: Session started with PTY") logrus.Info("SSH: Session started with PTY")
return nil return nil