221 lines
5.0 KiB
Go
221 lines
5.0 KiB
Go
package syncengine
|
|
|
|
import (
|
|
"context"
|
|
"fmt"
|
|
"io"
|
|
"os/exec"
|
|
"path/filepath"
|
|
"strings"
|
|
)
|
|
|
|
type RsyncResult struct {
|
|
ExitCode int
|
|
Stdout string
|
|
Stderr string
|
|
Stats *RsyncStats
|
|
}
|
|
|
|
type RsyncRunner struct {
|
|
sshDir string
|
|
privKey string
|
|
}
|
|
|
|
type SyncPairConfig struct {
|
|
ID int64
|
|
Name string
|
|
SourceMachineID *int64
|
|
SourcePath string
|
|
DestMachineID *int64
|
|
DestPath string
|
|
Direction string
|
|
RsyncFlags string
|
|
ExcludePatterns []string
|
|
Source string
|
|
Dest string
|
|
}
|
|
|
|
func NewRsyncRunner(sshDir, privKey string) *RsyncRunner {
|
|
return &RsyncRunner{sshDir: sshDir, privKey: privKey}
|
|
}
|
|
|
|
func (r *RsyncRunner) Run(ctx context.Context, pair *SyncPairConfig, onLine func(stream string, line string)) (*RsyncResult, error) {
|
|
cmd := r.buildRsyncCmd(pair)
|
|
return r.runCmd(ctx, cmd, onLine)
|
|
}
|
|
|
|
func (r *RsyncRunner) buildRsyncCmd(pair *SyncPairConfig) *exec.Cmd {
|
|
args := r.buildArgs(pair)
|
|
cmd := exec.CommandContext(context.Background(), "rsync", args...)
|
|
if r.privKey != "" {
|
|
sshCmd := fmt.Sprintf("ssh -i %s -o StrictHostKeyChecking=accept-new -o UserKnownHostsFile=%s",
|
|
r.privKey, strings.TrimRight(r.sshDir, "/")+"/known_hosts")
|
|
cmd.Args = append([]string{"rsync", "-e", sshCmd}, args...)
|
|
}
|
|
return cmd
|
|
}
|
|
|
|
func (r *RsyncRunner) buildArgs(pair *SyncPairConfig) []string {
|
|
var args []string
|
|
|
|
flags := strings.Fields(pair.RsyncFlags)
|
|
args = append(args, flags...)
|
|
|
|
for _, pattern := range pair.ExcludePatterns {
|
|
args = append(args, "--exclude="+pattern)
|
|
}
|
|
|
|
if pair.Direction == "mirror" {
|
|
args = append(args, "--delete")
|
|
}
|
|
|
|
src := ensureDirSlash(pair.Source)
|
|
if pair.Direction == "pull" {
|
|
args = append(args, pair.Dest, src)
|
|
} else {
|
|
args = append(args, src, pair.Dest)
|
|
}
|
|
|
|
return args
|
|
}
|
|
|
|
// ensureDirSlash guarantees the source path is treated by rsync as a
|
|
// directory whose contents are copied, regardless of whether the user
|
|
// supplied a trailing slash. This avoids the common foot-gun where
|
|
// "rsync host:/path/series /dest/" creates /dest/series/<contents> nested
|
|
// inside an extra "series" subdirectory.
|
|
func ensureDirSlash(p string) string {
|
|
if strings.HasSuffix(p, "/") {
|
|
return p
|
|
}
|
|
return p + "/"
|
|
}
|
|
|
|
type MachineKeys struct {
|
|
Host string
|
|
Port int
|
|
SSHUser string
|
|
PrivKey string
|
|
}
|
|
|
|
func (r *RsyncRunner) RunRemote(ctx context.Context, pair *SyncPairConfig, src *MachineKeys, dst *MachineKeys, destKey string, onLine func(stream string, line string)) (*RsyncResult, error) {
|
|
if src.Port == 0 {
|
|
src.Port = 22
|
|
}
|
|
if src.PrivKey == "" {
|
|
src.PrivKey = filepath.Join(r.sshDir, "id_ed25519")
|
|
}
|
|
if destKey == "" {
|
|
destKey = filepath.Join(r.sshDir, "id_ed25519")
|
|
}
|
|
|
|
args := r.buildArgs(pair)
|
|
|
|
innerSSH := fmt.Sprintf("ssh -i %s -o StrictHostKeyChecking=accept-new -o UserKnownHostsFile=%s",
|
|
destKey, filepath.Join(r.sshDir, "known_hosts"))
|
|
rsyncFlags := strings.Join(args[:len(args)-2], " ")
|
|
sourcePath := args[len(args)-2]
|
|
destPath := args[len(args)-1]
|
|
|
|
remoteCmd := fmt.Sprintf("rsync %s -e %q %s %s",
|
|
rsyncFlags, innerSSH, sourcePath, destPath)
|
|
|
|
sshArgs := []string{
|
|
"-i", src.PrivKey,
|
|
"-o", "StrictHostKeyChecking=accept-new",
|
|
"-o", "UserKnownHostsFile=" + filepath.Join(r.sshDir, "known_hosts"),
|
|
"-p", fmt.Sprintf("%d", src.Port),
|
|
fmt.Sprintf("%s@%s", src.SSHUser, src.Host),
|
|
}
|
|
sshArgs = append(sshArgs, remoteCmd)
|
|
|
|
cmd := exec.CommandContext(ctx, "ssh", sshArgs...)
|
|
return r.runCmd(ctx, cmd, onLine)
|
|
}
|
|
|
|
func (r *RsyncRunner) runCmd(ctx context.Context, cmd *exec.Cmd, onLine func(stream string, line string)) (*RsyncResult, error) {
|
|
stdout, err := cmd.StdoutPipe()
|
|
if err != nil {
|
|
return nil, fmt.Errorf("stdout pipe: %w", err)
|
|
}
|
|
stderr, err := cmd.StderrPipe()
|
|
if err != nil {
|
|
return nil, fmt.Errorf("stderr pipe: %w", err)
|
|
}
|
|
|
|
if err := cmd.Start(); err != nil {
|
|
return nil, fmt.Errorf("starting rsync: %w", err)
|
|
}
|
|
|
|
var outLines, errLines []string
|
|
done := make(chan struct{})
|
|
go func() {
|
|
br := io.Reader(stdout)
|
|
buf := make([]byte, 4096)
|
|
for {
|
|
n, err := br.Read(buf)
|
|
if n > 0 {
|
|
line := strings.TrimRight(string(buf[:n]), "\r\n")
|
|
if line != "" {
|
|
outLines = append(outLines, line)
|
|
if onLine != nil {
|
|
onLine("stdout", line)
|
|
}
|
|
}
|
|
}
|
|
if err != nil {
|
|
break
|
|
}
|
|
}
|
|
select {
|
|
case <-done:
|
|
default:
|
|
close(done)
|
|
}
|
|
}()
|
|
|
|
go func() {
|
|
br := io.Reader(stderr)
|
|
buf := make([]byte, 4096)
|
|
for {
|
|
n, err := br.Read(buf)
|
|
if n > 0 {
|
|
line := strings.TrimRight(string(buf[:n]), "\r\n")
|
|
if line != "" {
|
|
errLines = append(errLines, line)
|
|
if onLine != nil {
|
|
onLine("stderr", line)
|
|
}
|
|
}
|
|
}
|
|
if err != nil {
|
|
break
|
|
}
|
|
}
|
|
select {
|
|
case <-done:
|
|
default:
|
|
close(done)
|
|
}
|
|
}()
|
|
|
|
err = cmd.Wait()
|
|
<-done
|
|
|
|
result := &RsyncResult{
|
|
ExitCode: 0,
|
|
Stdout: strings.Join(outLines, "\n"),
|
|
Stderr: strings.Join(errLines, "\n"),
|
|
Stats: ParseFinalStats(strings.Join(outLines, "\n")),
|
|
}
|
|
if err != nil {
|
|
if exitErr, ok := err.(*exec.ExitError); ok {
|
|
result.ExitCode = exitErr.ExitCode()
|
|
} else {
|
|
result.ExitCode = -1
|
|
}
|
|
}
|
|
|
|
return result, nil
|
|
}
|