Let operators manage ordered action groups on each host, run them manually or on a schedule over SSH, and inspect disk-backed JSON run logs from a node detail UI.
229 lines
5.2 KiB
Go
229 lines
5.2 KiB
Go
package auth
|
|
|
|
import (
|
|
"bytes"
|
|
"fmt"
|
|
"net"
|
|
"strings"
|
|
"time"
|
|
|
|
"golang.org/x/crypto/ssh"
|
|
)
|
|
|
|
const defaultSSHRunTimeout = 5 * time.Minute
|
|
|
|
// SSHRunResult captures stdout, stderr, and exit status from a remote command.
|
|
type SSHRunResult struct {
|
|
Stdout string
|
|
Stderr string
|
|
ExitCode int
|
|
}
|
|
|
|
// RunSSHShell runs a single shell command on the remote host over SSH.
|
|
func RunSSHShell(
|
|
hostIP string,
|
|
username string,
|
|
privateKeyPEM string,
|
|
passphrase string,
|
|
command string,
|
|
env map[string]string,
|
|
timeout time.Duration,
|
|
) (SSHRunResult, error) {
|
|
if strings.TrimSpace(command) == "" {
|
|
return SSHRunResult{}, fmt.Errorf("command is required")
|
|
}
|
|
fullCommand := prependEnvExports(env, command)
|
|
return runSSHSession(hostIP, username, privateKeyPEM, passphrase, fullCommand, timeout)
|
|
}
|
|
|
|
// RunSSHScript uploads and runs a multi-line script via bash -s on the remote host.
|
|
func RunSSHScript(
|
|
hostIP string,
|
|
username string,
|
|
privateKeyPEM string,
|
|
passphrase string,
|
|
scriptBody string,
|
|
env map[string]string,
|
|
timeout time.Duration,
|
|
) (SSHRunResult, error) {
|
|
if strings.TrimSpace(scriptBody) == "" {
|
|
return SSHRunResult{}, fmt.Errorf("script body is required")
|
|
}
|
|
|
|
hostIP = strings.TrimSpace(hostIP)
|
|
username = strings.TrimSpace(username)
|
|
if hostIP == "" {
|
|
return SSHRunResult{}, fmt.Errorf("host IP is required")
|
|
}
|
|
if username == "" {
|
|
return SSHRunResult{}, fmt.Errorf("username is required")
|
|
}
|
|
if timeout <= 0 {
|
|
timeout = defaultSSHRunTimeout
|
|
}
|
|
|
|
signer, err := parseSigner(privateKeyPEM, passphrase)
|
|
if err != nil {
|
|
return SSHRunResult{}, err
|
|
}
|
|
|
|
config := &ssh.ClientConfig{
|
|
User: username,
|
|
Auth: []ssh.AuthMethod{
|
|
ssh.PublicKeys(signer),
|
|
},
|
|
HostKeyCallback: ssh.InsecureIgnoreHostKey(),
|
|
Timeout: sshDialTimeout,
|
|
}
|
|
|
|
address := net.JoinHostPort(hostIP, defaultSSHPort)
|
|
client, err := ssh.Dial("tcp", address, config)
|
|
if err != nil {
|
|
return SSHRunResult{}, fmt.Errorf("ssh dial %s: %w", address, err)
|
|
}
|
|
defer client.Close()
|
|
|
|
session, err := client.NewSession()
|
|
if err != nil {
|
|
return SSHRunResult{}, fmt.Errorf("ssh session: %w", err)
|
|
}
|
|
defer session.Close()
|
|
|
|
var stdoutBuf bytes.Buffer
|
|
var stderrBuf bytes.Buffer
|
|
session.Stdout = &stdoutBuf
|
|
session.Stderr = &stderrBuf
|
|
session.Stdin = strings.NewReader(scriptBody)
|
|
|
|
remoteCommand := prependEnvExports(env, "bash -s")
|
|
done := make(chan error, 1)
|
|
go func() {
|
|
done <- session.Run(remoteCommand)
|
|
}()
|
|
|
|
timer := time.NewTimer(timeout)
|
|
defer timer.Stop()
|
|
|
|
select {
|
|
case err := <-done:
|
|
return finishSSHResult(stdoutBuf.String(), stderrBuf.String(), err)
|
|
case <-timer.C:
|
|
_ = session.Close()
|
|
return SSHRunResult{
|
|
Stdout: stdoutBuf.String(),
|
|
Stderr: stderrBuf.String(),
|
|
ExitCode: -1,
|
|
}, fmt.Errorf("ssh command timed out after %s", timeout)
|
|
}
|
|
}
|
|
|
|
func runSSHSession(
|
|
hostIP string,
|
|
username string,
|
|
privateKeyPEM string,
|
|
passphrase string,
|
|
command string,
|
|
timeout time.Duration,
|
|
) (SSHRunResult, error) {
|
|
hostIP = strings.TrimSpace(hostIP)
|
|
username = strings.TrimSpace(username)
|
|
if hostIP == "" {
|
|
return SSHRunResult{}, fmt.Errorf("host IP is required")
|
|
}
|
|
if username == "" {
|
|
return SSHRunResult{}, fmt.Errorf("username is required")
|
|
}
|
|
if timeout <= 0 {
|
|
timeout = defaultSSHRunTimeout
|
|
}
|
|
|
|
signer, err := parseSigner(privateKeyPEM, passphrase)
|
|
if err != nil {
|
|
return SSHRunResult{}, err
|
|
}
|
|
|
|
config := &ssh.ClientConfig{
|
|
User: username,
|
|
Auth: []ssh.AuthMethod{
|
|
ssh.PublicKeys(signer),
|
|
},
|
|
HostKeyCallback: ssh.InsecureIgnoreHostKey(),
|
|
Timeout: sshDialTimeout,
|
|
}
|
|
|
|
address := net.JoinHostPort(hostIP, defaultSSHPort)
|
|
client, err := ssh.Dial("tcp", address, config)
|
|
if err != nil {
|
|
return SSHRunResult{}, fmt.Errorf("ssh dial %s: %w", address, err)
|
|
}
|
|
defer client.Close()
|
|
|
|
session, err := client.NewSession()
|
|
if err != nil {
|
|
return SSHRunResult{}, fmt.Errorf("ssh session: %w", err)
|
|
}
|
|
defer session.Close()
|
|
|
|
var stdoutBuf bytes.Buffer
|
|
var stderrBuf bytes.Buffer
|
|
session.Stdout = &stdoutBuf
|
|
session.Stderr = &stderrBuf
|
|
|
|
done := make(chan error, 1)
|
|
go func() {
|
|
done <- session.Run(command)
|
|
}()
|
|
|
|
timer := time.NewTimer(timeout)
|
|
defer timer.Stop()
|
|
|
|
select {
|
|
case err := <-done:
|
|
return finishSSHResult(stdoutBuf.String(), stderrBuf.String(), err)
|
|
case <-timer.C:
|
|
_ = session.Close()
|
|
return SSHRunResult{
|
|
Stdout: stdoutBuf.String(),
|
|
Stderr: stderrBuf.String(),
|
|
ExitCode: -1,
|
|
}, fmt.Errorf("ssh command timed out after %s", timeout)
|
|
}
|
|
}
|
|
|
|
func finishSSHResult(stdout string, stderr string, err error) (SSHRunResult, error) {
|
|
result := SSHRunResult{
|
|
Stdout: stdout,
|
|
Stderr: stderr,
|
|
ExitCode: 0,
|
|
}
|
|
if err == nil {
|
|
return result, nil
|
|
}
|
|
if exitErr, ok := err.(*ssh.ExitError); ok {
|
|
result.ExitCode = exitErr.ExitStatus()
|
|
return result, nil
|
|
}
|
|
result.ExitCode = -1
|
|
return result, err
|
|
}
|
|
|
|
func prependEnvExports(env map[string]string, command string) string {
|
|
if len(env) == 0 {
|
|
return command
|
|
}
|
|
var builder strings.Builder
|
|
for name, value := range env {
|
|
builder.WriteString("export ")
|
|
builder.WriteString(shellQuote(name))
|
|
builder.WriteString("=")
|
|
builder.WriteString(shellQuote(value))
|
|
builder.WriteString("; ")
|
|
}
|
|
builder.WriteString(command)
|
|
return builder.String()
|
|
}
|
|
|
|
func shellQuote(value string) string {
|
|
return "'" + strings.ReplaceAll(value, "'", `'\''`) + "'"
|
|
}
|