Initial commit: forge-tools-ssh — MCP-сервер для администрирования по SSH
This commit is contained in:
@@ -0,0 +1,396 @@
|
||||
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)
|
||||
}
|
||||
Reference in New Issue
Block a user