package sshmanager import ( "bytes" "context" "crypto/sha256" "encoding/base64" "fmt" "net" "os" "strings" "time" "golang.org/x/crypto/ssh" ) type ConnResult struct { Success bool Output string Error string Fingerprint string } func dialSSH(ctx context.Context, host string, port int, user, privKeyPath, knownHostsPath string, strictHostKeyChecking bool) (*ssh.Client, string, ssh.PublicKey, error) { addr := fmt.Sprintf("%s:%d", host, port) auths := []ssh.AuthMethod{} if privKeyPath != "" { key, err := os.ReadFile(privKeyPath) if err != nil { return nil, "", nil, fmt.Errorf("reading private key: %w", err) } signer, err := ssh.ParsePrivateKey(key) if err != nil { return nil, "", nil, fmt.Errorf("parsing private key: %w", err) } auths = append(auths, ssh.PublicKeys(signer)) } var capturedFingerprint string var capturedPubKey ssh.PublicKey callback, err := NewKnownHostsCallback(knownHostsPath, strictHostKeyChecking) if err != nil { return nil, "", nil, fmt.Errorf("creating host key callback: %w", err) } hostKeyCallback := func(hostname string, remote net.Addr, key ssh.PublicKey) error { h := sha256.Sum256(key.Marshal()) capturedFingerprint = "SHA256:" + base64.RawStdEncoding.EncodeToString(h[:]) capturedPubKey = key return callback(hostname, remote, key) } cfg := &ssh.ClientConfig{ User: user, Auth: auths, HostKeyCallback: hostKeyCallback, Timeout: 10 * time.Second, } ctx, cancel := context.WithTimeout(ctx, 30*time.Second) defer cancel() conn, err := ssh.Dial("tcp", addr, cfg) if err != nil { if strings.Contains(err.Error(), "known_hosts") || strings.Contains(err.Error(), "host key") { return nil, capturedFingerprint, capturedPubKey, fmt.Errorf("host key verification failed: %v", err) } return nil, capturedFingerprint, capturedPubKey, fmt.Errorf("connection failed: %v", err) } return conn, capturedFingerprint, capturedPubKey, nil } func TestSSHConnection(ctx context.Context, host string, port int, user, privKeyPath, knownHostsPath string, strictHostKeyChecking bool) (*ConnResult, error) { conn, fingerprint, _, err := dialSSH(ctx, host, port, user, privKeyPath, knownHostsPath, strictHostKeyChecking) if err != nil { return &ConnResult{ Success: false, Error: err.Error(), Fingerprint: fingerprint, }, nil } defer conn.Close() session, err := conn.NewSession() if err != nil { return &ConnResult{Success: false, Error: fmt.Sprintf("session: %v", err), Fingerprint: fingerprint}, nil } defer session.Close() var stdout, stderr bytes.Buffer session.Stdout = &stdout session.Stderr = &stderr if err := session.Run("echo ok && uname -a"); err != nil { return &ConnResult{ Success: false, Error: fmt.Sprintf("exec: %v, stderr: %s", err, stderr.String()), Fingerprint: fingerprint, }, nil } return &ConnResult{ Success: true, Output: stdout.String(), Fingerprint: fingerprint, }, nil } func ConnectForApproval(ctx context.Context, host string, port int, user, privKeyPath, knownHostsPath string) (*ssh.Client, string, ssh.PublicKey, error) { return dialSSH(ctx, host, port, user, privKeyPath, knownHostsPath, false) }