145 lines
5 KiB
Go
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
|
|
}
|