397 lines
11 KiB
Go
397 lines
11 KiB
Go
package ssh
|
|
|
|
import (
|
|
"context"
|
|
"errors"
|
|
"fmt"
|
|
"io"
|
|
"log"
|
|
"net"
|
|
"strings"
|
|
"sync"
|
|
"time"
|
|
|
|
"github.com/pkg/sftp"
|
|
"golang.org/x/crypto/ssh"
|
|
)
|
|
|
|
// Client represents a single SSH connection with state tracking.
|
|
type Client struct {
|
|
alias string
|
|
conn *ssh.Client
|
|
sftp *sftp.Client
|
|
cwd string
|
|
mu sync.Mutex
|
|
creds Credentials
|
|
}
|
|
|
|
// Credentials holds SSH connection parameters.
|
|
type Credentials struct {
|
|
Host string
|
|
Port int
|
|
Username string
|
|
Password string
|
|
PrivateKey ssh.Signer
|
|
Via string
|
|
|
|
// Target is an optional destination host routed through a PAM gateway.
|
|
// When set, the SSH username sent to the gateway becomes
|
|
// "Username@Target" and authentication switches to multi-stage
|
|
// keyboard-interactive.
|
|
Target string
|
|
// TargetPassword is the password of the Target server account,
|
|
// requested by the PAM gateway as a second authentication stage.
|
|
TargetPassword string
|
|
}
|
|
|
|
// NewClient creates a new SSH client.
|
|
func NewClient(ctx context.Context, alias string, creds Credentials, jumpClient *Client) (*Client, error) {
|
|
client := &Client{
|
|
alias: alias,
|
|
creds: creds,
|
|
cwd: "",
|
|
}
|
|
|
|
if err := client.connect(ctx, jumpClient); err != nil {
|
|
return nil, err
|
|
}
|
|
|
|
return client, nil
|
|
}
|
|
|
|
// connect establishes the SSH connection. ctx отменяем: уважается и при
|
|
// TCP-dial, и при пробе «pwd» (иначе полу-готовый хост, принявший TCP, но
|
|
// не отдающий шелл, мог бы подвесить вызов без срока — см. исходный
|
|
// зависонный кейс).
|
|
func (c *Client) connect(ctx context.Context, jumpClient *Client) error {
|
|
c.mu.Lock()
|
|
defer c.mu.Unlock()
|
|
|
|
if ctx == nil {
|
|
ctx = context.Background()
|
|
}
|
|
|
|
// Close stale SFTP client before closing the connection
|
|
if c.sftp != nil {
|
|
c.sftp.Close()
|
|
c.sftp = nil
|
|
}
|
|
|
|
if c.conn != nil {
|
|
c.conn.Close()
|
|
c.conn = nil
|
|
}
|
|
|
|
// PAM gateway routing: the gateway parses "user@target" from the SSH
|
|
// username and proxies the session to the target host. The plain
|
|
// username is kept in Credentials for reconnect and logging.
|
|
user := c.creds.Username
|
|
if c.creds.Target != "" {
|
|
user = fmt.Sprintf("%s@%s", c.creds.Username, c.creds.Target)
|
|
}
|
|
|
|
config := &ssh.ClientConfig{
|
|
User: user,
|
|
HostKeyCallback: ssh.InsecureIgnoreHostKey(),
|
|
Timeout: 30 * time.Second,
|
|
}
|
|
|
|
var authMethods []ssh.AuthMethod
|
|
if c.creds.PrivateKey != nil {
|
|
authMethods = append(authMethods, ssh.PublicKeys(c.creds.PrivateKey))
|
|
}
|
|
if c.creds.Target != "" {
|
|
// SafeInspect PAM authenticates in two stages:
|
|
// 1. keyboard-interactive with the PAM gateway password (the server
|
|
// replies with partial success),
|
|
// 2. the standard "password" method with the target server password.
|
|
// The Go SSH client tries each AuthMethod in order on partial success,
|
|
// so we supply both methods in the order above.
|
|
authMethods = append(authMethods, pamKeyboardInteractive(c.creds))
|
|
if c.creds.TargetPassword != "" {
|
|
authMethods = append(authMethods, ssh.Password(c.creds.TargetPassword))
|
|
}
|
|
} else if c.creds.Password != "" {
|
|
authMethods = append(authMethods, ssh.Password(c.creds.Password))
|
|
}
|
|
if len(authMethods) == 0 {
|
|
return errors.New("no authentication method provided (key or password required)")
|
|
}
|
|
config.Auth = authMethods
|
|
|
|
addr := fmt.Sprintf("%s:%d", c.creds.Host, c.creds.Port)
|
|
|
|
var conn *ssh.Client
|
|
var err error
|
|
|
|
if jumpClient != nil {
|
|
jumpConn := jumpClient.conn
|
|
if jumpConn == nil {
|
|
return errors.New("jump host not connected")
|
|
}
|
|
|
|
netConn, err := dialWithCtx(ctx, func() (net.Conn, error) {
|
|
return jumpConn.Dial("tcp", addr)
|
|
})
|
|
if err != nil {
|
|
return fmt.Errorf("failed to dial through jump host: %w", err)
|
|
}
|
|
|
|
ncc, chans, reqs, err := ssh.NewClientConn(netConn, addr, config)
|
|
if err != nil {
|
|
netConn.Close()
|
|
return fmt.Errorf("failed to create client connection through jump: %w", err)
|
|
}
|
|
|
|
conn = ssh.NewClient(ncc, chans, reqs)
|
|
} else {
|
|
// net.Dialer.DialContext вместо ssh.Dial: уважает ctx (отмену
|
|
// задачи), сохраняя 30-секундный лимит на установку соединения.
|
|
raw, err := (&net.Dialer{Timeout: 30 * time.Second}).DialContext(ctx, "tcp", addr)
|
|
if err != nil {
|
|
return fmt.Errorf("failed to connect: %w", err)
|
|
}
|
|
ncc, chans, reqs, err := ssh.NewClientConn(raw, addr, config)
|
|
if err != nil {
|
|
raw.Close()
|
|
return fmt.Errorf("failed to create client connection: %w", err)
|
|
}
|
|
conn = ssh.NewClient(ncc, chans, reqs)
|
|
}
|
|
|
|
c.conn = conn
|
|
|
|
// Проба «pwd» с дедлайном: соединение установлено, но хост может не
|
|
// отдать шелл - ограничиваем ожидание, чтобы не подвесить вызов.
|
|
probeCtx, probeCancel := context.WithTimeout(ctx, 15*time.Second)
|
|
defer probeCancel()
|
|
output, err := c.runRaw(probeCtx, "pwd")
|
|
if err != nil {
|
|
c.cwd = "~"
|
|
} else {
|
|
c.cwd = strings.TrimSpace(output)
|
|
}
|
|
|
|
log.Printf("ssh: connected %s@%s (%s)", c.creds.Username, c.creds.Host, c.alias)
|
|
return nil
|
|
}
|
|
|
|
// dialWithCtx выполняет сетевой dial с уважением к отмене ctx: если ctx
|
|
// отменён раньше, чем dial вернул соединение, канал разрывается. jump-
|
|
// диал не имеет ctx-версии, поэтому оборачиваем в select.
|
|
func dialWithCtx(ctx context.Context, dial func() (net.Conn, error)) (net.Conn, error) {
|
|
type res struct {
|
|
conn net.Conn
|
|
err error
|
|
}
|
|
ch := make(chan res, 1)
|
|
go func() {
|
|
c, err := dial()
|
|
ch <- res{c, err}
|
|
}()
|
|
select {
|
|
case <-ctx.Done():
|
|
return nil, ctx.Err()
|
|
case r := <-ch:
|
|
return r.conn, r.err
|
|
}
|
|
}
|
|
|
|
// pamKeyboardInteractive returns an ssh.AuthMethod for the first stage of a
|
|
// SafeInspect PAM gateway. The gateway presents a keyboard-interactive
|
|
// challenge asking for the PAM user's password; after the server signals
|
|
// partial success, the target server password is supplied by the following
|
|
// "password" AuthMethod. Challenge rounds that carry instructions only (no
|
|
// questions) are answered with an empty response.
|
|
func pamKeyboardInteractive(creds Credentials) ssh.AuthMethod {
|
|
return ssh.KeyboardInteractive(func(user, instruction string, questions []string, echos []bool) ([]string, error) {
|
|
if len(questions) == 0 {
|
|
return []string{}, nil
|
|
}
|
|
answers := make([]string, len(questions))
|
|
for i := range answers {
|
|
answers[i] = creds.Password
|
|
}
|
|
return answers, nil
|
|
})
|
|
}
|
|
|
|
// runRaw executes a command without CWD handling. Команда ограничена ctx
|
|
// (дедлайном): уважается при отмене задачи/превышении лимита, иначе
|
|
// полу-ответивший хост мог бы подвесить соединение без срока.
|
|
func (c *Client) runRaw(ctx context.Context, cmd string) (string, error) {
|
|
session, err := c.conn.NewSession()
|
|
if err != nil {
|
|
return "", fmt.Errorf("failed to create session: %w", err)
|
|
}
|
|
defer session.Close()
|
|
|
|
type out struct {
|
|
s string
|
|
err error
|
|
}
|
|
ch := make(chan out, 1)
|
|
go func() {
|
|
o, e := session.CombinedOutput(cmd)
|
|
ch <- out{string(o), e}
|
|
}()
|
|
|
|
select {
|
|
case <-ctx.Done():
|
|
_ = session.Close() // разрывает блокирующий вызов в горутине
|
|
return "", ctx.Err()
|
|
case r := <-ch:
|
|
return r.s, r.err
|
|
}
|
|
}
|
|
|
|
// Run executes a command with CWD tracking.
|
|
func (c *Client) Run(ctx context.Context, cmd string) (*RunResult, error) {
|
|
c.mu.Lock()
|
|
defer c.mu.Unlock()
|
|
|
|
if c.conn == nil {
|
|
return nil, errors.New("not connected")
|
|
}
|
|
|
|
delimiter := fmt.Sprintf("___MCP_PWD_%d___", time.Now().UnixNano())
|
|
wrappedCmd := fmt.Sprintf(
|
|
`cd %q && %s; __EXIT__=$?; echo ""; echo "%s"; pwd; exit $__EXIT__`,
|
|
c.cwd, strings.TrimRight(cmd, " \t\r\n"), delimiter,
|
|
)
|
|
|
|
session, err := c.conn.NewSession()
|
|
if err != nil {
|
|
return nil, fmt.Errorf("failed to create session: %w", err)
|
|
}
|
|
defer session.Close()
|
|
|
|
stdout, _ := session.StdoutPipe()
|
|
stderr, _ := session.StderrPipe()
|
|
|
|
if err := session.Start(wrappedCmd); err != nil {
|
|
return nil, fmt.Errorf("failed to start command: %w", err)
|
|
}
|
|
|
|
type readResult struct {
|
|
stdout []byte
|
|
stderr []byte
|
|
}
|
|
resultChan := make(chan readResult, 1)
|
|
|
|
const maxStdout = 10 * 1024 * 1024 // 10 MB
|
|
const maxStderr = 1 * 1024 * 1024 // 1 MB
|
|
go func() {
|
|
stdoutBytes, _ := io.ReadAll(io.LimitReader(stdout, maxStdout))
|
|
stderrBytes, _ := io.ReadAll(io.LimitReader(stderr, maxStderr))
|
|
resultChan <- readResult{stdout: stdoutBytes, stderr: stderrBytes}
|
|
}()
|
|
|
|
var res readResult
|
|
select {
|
|
case <-ctx.Done():
|
|
_ = session.Signal(ssh.SIGKILL)
|
|
_ = session.Close() // Unblock io.ReadAll by closing pipes
|
|
// Wait briefly for the reader goroutine to finish
|
|
select {
|
|
case <-resultChan:
|
|
case <-time.After(2 * time.Second):
|
|
}
|
|
return nil, ctx.Err()
|
|
case res = <-resultChan:
|
|
}
|
|
|
|
var exitCode int
|
|
if err := session.Wait(); err != nil {
|
|
if exitErr, ok := err.(*ssh.ExitError); ok {
|
|
exitCode = exitErr.ExitStatus()
|
|
} else {
|
|
return nil, fmt.Errorf("command failed: %w", err)
|
|
}
|
|
}
|
|
|
|
stdoutStr := string(res.stdout)
|
|
cleanOutput := stdoutStr
|
|
if idx := strings.Index(stdoutStr, delimiter); idx != -1 {
|
|
cleanOutput = stdoutStr[:idx]
|
|
remaining := strings.TrimSpace(stdoutStr[idx+len(delimiter):])
|
|
if remaining != "" {
|
|
c.cwd = remaining
|
|
}
|
|
}
|
|
|
|
return &RunResult{
|
|
Stdout: strings.TrimSpace(cleanOutput),
|
|
Stderr: strings.TrimSpace(string(res.stderr)),
|
|
ExitCode: exitCode,
|
|
CWD: c.cwd,
|
|
}, nil
|
|
}
|
|
|
|
// RunResult contains command execution result.
|
|
type RunResult struct {
|
|
Stdout string
|
|
Stderr string
|
|
ExitCode int
|
|
CWD string
|
|
}
|
|
|
|
// SFTP returns the SFTP client.
|
|
func (c *Client) SFTP() (*sftp.Client, error) {
|
|
c.mu.Lock()
|
|
defer c.mu.Unlock()
|
|
|
|
if c.conn == nil {
|
|
return nil, errors.New("not connected")
|
|
}
|
|
|
|
if c.sftp != nil {
|
|
return c.sftp, nil
|
|
}
|
|
|
|
sftpClient, err := sftp.NewClient(c.conn)
|
|
if err != nil {
|
|
return nil, fmt.Errorf("failed to create SFTP client: %w", err)
|
|
}
|
|
|
|
c.sftp = sftpClient
|
|
return c.sftp, nil
|
|
}
|
|
|
|
// Close closes the connection.
|
|
func (c *Client) Close() error {
|
|
c.mu.Lock()
|
|
defer c.mu.Unlock()
|
|
|
|
if c.sftp != nil {
|
|
c.sftp.Close()
|
|
c.sftp = nil
|
|
}
|
|
|
|
if c.conn != nil {
|
|
err := c.conn.Close()
|
|
c.conn = nil
|
|
return err
|
|
}
|
|
|
|
return nil
|
|
}
|
|
|
|
// Alias returns the connection alias.
|
|
func (c *Client) Alias() string {
|
|
return c.alias
|
|
}
|
|
|
|
// CWD returns the current working directory.
|
|
func (c *Client) CWD() string {
|
|
c.mu.Lock()
|
|
defer c.mu.Unlock()
|
|
return c.cwd
|
|
}
|
|
|
|
// Reconnect attempts to reconnect.
|
|
func (c *Client) Reconnect(ctx context.Context, jumpClient *Client) error {
|
|
log.Printf("ssh: reconnecting %s", c.alias)
|
|
return c.connect(ctx, jumpClient)
|
|
}
|