[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
parent 16f7e1e148
commit 2da4512272
5 changed files with 100 additions and 103 deletions

View File

@@ -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

View File

@@ -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

View File

@@ -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() {

73
client/ssh/session.go Normal file
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
}

View File

@@ -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