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