gapmind/services/orchestra/service/systemd.go
2026-09-11 23:59:05 -04:00

145 lines
5 KiB
Go

package service
import (
"context"
"encoding/base64"
"fmt"
"os"
"os/exec"
"path"
"path/filepath"
"sort"
"strings"
"time"
)
const defaultCommandTimeout = 45 * time.Second
type SSHClient struct {
remote string
askPassPath string
password string
sshConfig string
timeout time.Duration
}
func NewSSHClient(remote, password string) (*SSHClient, error) {
if remote == "" {
return nil, fmt.Errorf("remote host is required")
}
sshConfig := "/dev/null"
if home, err := os.UserHomeDir(); err == nil {
candidate := filepath.Join(home, ".ssh", "config")
if info, statErr := os.Stat(candidate); statErr == nil && info.Mode().IsRegular() {
sshConfig = candidate
}
}
client := &SSHClient{remote: remote, password: password, sshConfig: sshConfig, timeout: defaultCommandTimeout}
if password == "" {
return client, nil
}
file, err := os.CreateTemp("", "gapmind-ssh-askpass-*")
if err != nil {
return nil, fmt.Errorf("create SSH password helper: %w", err)
}
client.askPassPath = file.Name()
if _, err := file.WriteString("#!/bin/sh\nprintf '%s\\n' \"$GAPMIND_SSH_PASSWORD\"\n"); err != nil {
_ = file.Close()
client.Close()
return nil, err
}
if err := file.Close(); err != nil {
client.Close()
return nil, err
}
if err := os.Chmod(client.askPassPath, 0700); err != nil {
client.Close()
return nil, err
}
return client, nil
}
func (c *SSHClient) Close() {
if c.askPassPath != "" {
_ = os.Remove(c.askPassPath)
}
}
func (c *SSHClient) SetTimeout(timeout time.Duration) {
if timeout > 0 {
c.timeout = timeout
}
}
func (c *SSHClient) Remote() string {
return c.remote
}
func (c *SSHClient) command(ctx context.Context, name string, args ...string) *exec.Cmd {
connectionOptions := []string{
"-F", c.sshConfig,
"-o", "ConnectTimeout=12",
"-o", "ConnectionAttempts=1",
"-o", "ServerAliveInterval=10",
"-o", "ServerAliveCountMax=2",
"-o", "NumberOfPasswordPrompts=1",
"-o", "StrictHostKeyChecking=accept-new",
"-o", "ClearAllForwardings=yes",
}
args = append(connectionOptions, args...)
if c.password != "" {
args = append([]string{"-o", "BatchMode=no", "-o", "PreferredAuthentications=password,keyboard-interactive", "-o", "PubkeyAuthentication=no"}, args...)
}
cmd := exec.CommandContext(ctx, name, args...)
if c.password != "" {
cmd.Env = append(os.Environ(), "SSH_ASKPASS="+c.askPassPath, "SSH_ASKPASS_REQUIRE=force", "DISPLAY=gapmind:0", "GAPMIND_SSH_PASSWORD="+c.password)
}
return cmd
}
func (c *SSHClient) Run(name string, args ...string) ([]byte, error) {
ctx, cancel := context.WithTimeout(context.Background(), c.timeout)
defer cancel()
output, err := c.command(ctx, name, args...).CombinedOutput()
if ctx.Err() == context.DeadlineExceeded {
return output, fmt.Errorf("%s timed out after %s", name, c.timeout)
}
return output, err
}
func (u UserUnit) Render() string {
lines := []string{"[Unit]", "Description=" + u.Description, "After=network-online.target", "", "[Service]", "Type=simple", "Restart=on-failure", "RestartSec=5"}
keys := make([]string, 0, len(u.Environment))
for key := range u.Environment {
keys = append(keys, key)
}
sort.Strings(keys)
for _, key := range keys {
if value := u.Environment[key]; value != "" {
lines = append(lines, "Environment="+key+"="+systemdQuote(value))
}
}
lines = append(lines, "ExecStart="+u.ExecStart, "", "[Install]", "WantedBy=default.target", "")
return strings.Join(lines, "\n")
}
func systemdQuote(value string) string {
return `"` + strings.NewReplacer(`\`, `\\`, `"`, `\"`, "\n", "").Replace(value) + `"`
}
func CopyAndInstall(client *SSHClient, localBinary, remoteBinary, unitName string, unit UserUnit) error {
if output, err := client.Run("ssh", client.remote, "mkdir", "-p", path.Dir(remoteBinary), ".config/systemd/user"); err != nil {
return fmt.Errorf("create remote directories: %w: %s", err, output)
}
temporaryBinary := remoteBinary + ".upload"
if output, err := client.Run("scp", localBinary, client.remote+":"+temporaryBinary); err != nil {
return fmt.Errorf("copy binary: %w: %s", err, output)
}
encoded := base64.StdEncoding.EncodeToString([]byte(unit.Render()))
remoteUnit := ".config/systemd/user/" + unitName
command := fmt.Sprintf("set -eu; user=$(id -un); loginctl enable-linger $user >/dev/null 2>&1 || true; runtime=/run/user/$(id -u); export XDG_RUNTIME_DIR=$runtime DBUS_SESSION_BUS_ADDRESS=unix:path=$runtime/bus; for attempt in 1 2 3 4 5; do test -S $runtime/bus && break; sleep 1; done; test -S $runtime/bus || { echo 'systemd user bus is unavailable after enabling lingering' >&2; exit 1; }; test ! -d %q || { echo 'deployment target is a directory: %s' >&2; exit 1; }; chmod +x %q; mv -f %q %q; printf %%s %q | base64 -d > %q; systemctl --user daemon-reload; systemctl --user enable --now %q; systemctl --user --no-pager is-active %q", remoteBinary, remoteBinary, temporaryBinary, temporaryBinary, remoteBinary, encoded, remoteUnit, unitName, unitName)
if output, err := client.Run("ssh", client.remote, "sh", "-lc", command); err != nil {
return fmt.Errorf("install systemd unit: %w: %s", err, output)
}
return nil
}