Files

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)
}