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") } if pair.Direction == "pull" { args = append(args, pair.Dest, pair.Source) } else { args = append(args, pair.Source, pair.Dest) } return args } 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") } srcUserHost := fmt.Sprintf("%s@%s", src.SSHUser, src.Host) 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], " ") var sourcePath, destPath string if pair.Direction == "pull" { sourcePath = args[len(args)-1] destPath = args[len(args)-2] } else { sourcePath = args[len(args)-2] destPath = args[len(args)-1] } var remoteSrc, remoteDst string if pair.Direction == "pull" { remoteSrc = sourcePath remoteDst = destPath if !strings.Contains(destPath, "@") { remoteDst = fmt.Sprintf("%s@%s:%s", dst.SSHUser, dst.Host, destPath) } } else { remoteDst = destPath if !strings.Contains(destPath, "@") { remoteDst = fmt.Sprintf("%s@%s:%s", dst.SSHUser, dst.Host, destPath) } remoteSrc = fmt.Sprintf("%s:%s", srcUserHost, sourcePath) } remoteCmd := fmt.Sprintf("rsync %s -e %q %s %s", rsyncFlags, innerSSH, remoteSrc, remoteDst) 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 }