mirror of
https://github.com/netbirdio/netbird.git
synced 2026-10-09 06:59:08 +02:00
[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:
co-authored by
Claude Fable 5
parent
16f7e1e148
commit
2da4512272
@@ -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
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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() {
|
||||||
|
|||||||
@@ -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
|
||||||
|
}
|
||||||
@@ -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
|
||||||
|
|||||||
Reference in New Issue
Block a user