package main import ( "encoding/json" "errors" "flag" "fmt" "io" "os" "os/exec" "strconv" "sync" "sync/atomic" "time" ) func cmdAgent(args []string) error { fs := flag.NewFlagSet("agent", flag.ContinueOnError) role := fs.String("role", "", "control|sink|source-stream|fs-send|fs-recv (internal)") path := fs.String("path", "", "path to read/write") size := fs.Int64("size", 0, "total sync size in bytes") base := fs.Int64("base", 0, "byte offset the window starts at (clone-disk boot region / offset partition)") blockSize := fs.Int64("block-size", defaultBlockSize, "block size in bytes") fsType := fs.String("fs", "", "filesystem type (fs-send/fs-recv)") fsTool := fs.String("fstool", "", "fs-image tool family (fs-send/fs-recv)") peerDisk := fs.String("peerdisk", "", "this helper's whole-disk/image path (fs-send/fs-recv)") shrinkTo := fs.Int64("shrink", 0, "shrink this fs to N bytes before sending (fs-send)") if err := fs.Parse(args); err != nil { return err } switch *role { case "control": return runControlAgent() case "sink": return runSinkRole(*path, *base, *size, *blockSize) case "source-stream": return runSourceStreamRole(*path, *base, *size, *blockSize) case "fs-send": n, _ := strconv.Atoi(*path) return runFSSendRole(n, *fsType, *fsTool, *peerDisk, *shrinkTo) case "fs-recv": n, _ := strconv.Atoi(*path) return runFSRecvRole(n, *fsType, *fsTool, *peerDisk) default: return fmt.Errorf("agent: unknown or missing --role %q", *role) } } // --------------------------------------------------------------------- // control role: long-lived per-side orchestration agent, driven by the // manager over stdin/stdout with CtrlMsg frames. // --------------------------------------------------------------------- func runControlAgent() error { in := NewFrameReader(os.Stdin) out := NewFrameWriter(os.Stdout) for { typ, payload, err := in.ReadFrame() if err != nil { return nil // manager closed the pipe; nothing left to do } if typ != frameCtrlJSON { return fmt.Errorf("control agent: unexpected frame type %d", typ) } var m CtrlMsg if err := json.Unmarshal(payload, &m); err != nil { return fmt.Errorf("control agent: decode message: %w", err) } switch m.Type { case msgStat: handleStat(out, m) case msgPrepare: handlePrepare(out, m) case msgConnectPush: runPushDriver(m, out) case msgConnectPull: runPullDriver(m, out) case msgProbeDisk: handleProbeDisk(out, m) case msgBuildLayout: handleBuildLayout(out, m) case msgClonePartition: handleClonePartition(out, m) case msgReinstallBoot: handleReinstallBoot(out, m) case msgClose: detachAllDisks() _ = out.WriteJSON(CtrlMsg{Type: msgBye}) return nil default: _ = out.WriteJSON(CtrlMsg{Type: msgError, Message: fmt.Sprintf("unknown command %q", m.Type)}) } } } func handleStat(out *FrameWriter, m CtrlMsg) { info, err := statPath(m.Path) if err != nil { _ = out.WriteJSON(CtrlMsg{Type: msgError, Message: err.Error(), NeedPriv: isPermErr(err)}) return } _ = out.WriteJSON(CtrlMsg{Type: msgStatOK, Exists: info.Exists, IsDevice: info.IsDevice, Size: info.Size}) } func handlePrepare(out *FrameWriter, m CtrlMsg) { if err := prepareDest(m.Path, m.Size); err != nil { _ = out.WriteJSON(CtrlMsg{Type: msgError, Message: err.Error(), NeedPriv: isPermErr(err)}) return } _ = out.WriteJSON(CtrlMsg{Type: msgPrepareOK, Size: m.Size}) } // isPermErr reports whether err (or anything it wraps, including a bare // syscall errno from an ioctl) is a permission failure — the signal that // retrying the agent under sudo might help. func isPermErr(err error) bool { return errors.Is(err, os.ErrPermission) } // peerAgentCommand builds the *exec.Cmd for a one-shot helper (sink / // source-stream) on the peer — locally or over ssh — wrapping it in sudo // when req.Sudo is set (interactive sudo locally; "sudo -n" remotely, since // a password prompt would corrupt the binary stream on stdout). func peerAgentCommand(req CtrlMsg, tailArgs []string) *exec.Cmd { if req.PeerLocal { if req.Sudo { return sudoLocalCommand(tailArgs) } return localAgentCommand(tailArgs) } peerArgs := append([]string{req.RemoteBin}, tailArgs...) if req.Sudo { peerArgs = append([]string{"sudo", "-n", "--"}, peerArgs...) } return sshCommand(req.SSHBin, req.SSHOpts, true, req.ConnectTimeoutSec, req.PeerUser, req.PeerHost, peerArgs) } // --------------------------------------------------------------------- // push driver: runs inside the SOURCE control agent. Spawns ssh straight // to the destination host, running the "sink" role, and — once it answers // READY — performs the whole read/hash/compare/send loop itself. // --------------------------------------------------------------------- func runPushDriver(req CtrlMsg, out *FrameWriter) { tailArgs := []string{ "agent", "--role", "sink", "--path", req.PeerPath, "--base", strconv.FormatInt(req.PeerBase, 10), "--size", strconv.FormatInt(req.Size, 10), "--block-size", strconv.FormatInt(req.BlockSize, 10), } cmd := peerAgentCommand(req, tailArgs) stdin, err := cmd.StdinPipe() if err != nil { _ = out.WriteJSON(CtrlMsg{Type: msgPushFailed, Reason: err.Error()}) return } stdout, err := cmd.StdoutPipe() if err != nil { _ = out.WriteJSON(CtrlMsg{Type: msgPushFailed, Reason: err.Error()}) return } stderrBuf := newLimitedBuffer(4096) cmd.Stderr = stderrBuf if err := cmd.Start(); err != nil { _ = out.WriteJSON(CtrlMsg{Type: msgPushFailed, Reason: fmt.Sprintf("start ssh: %v", err)}) return } fw := NewFrameWriter(stdin) fr := NewFrameReader(stdout) timeout := time.Duration(req.ConnectTimeoutSec+2) * time.Second if err := waitReady(fr, timeout); err != nil { _ = cmd.Process.Kill() _ = cmd.Wait() _ = out.WriteJSON(CtrlMsg{Type: msgPushFailed, Reason: fmt.Sprintf("%v (remote stderr: %s)", err, stderrBuf.String())}) return } // Handshake succeeded: we're committed to push for this run. The sink // streams the destination's current per-block hashes as it scans; the // source loop consumes them in order and reads/compares its own blocks // as they arrive, so the two scans and the transfer all overlap. srcFile, err := os.Open(req.Path) if err != nil { _ = cmd.Process.Kill() _ = cmd.Wait() _ = out.WriteJSON(CtrlMsg{Type: msgError, Message: err.Error(), NeedPriv: isPermErr(err)}) return } defer srcFile.Close() if fatalErr := pumpPush(req, alignmentFor(req.Path), srcFile, fw, fr, out); fatalErr != nil { _ = cmd.Process.Kill() _ = cmd.Wait() _ = out.WriteJSON(CtrlMsg{Type: msgError, Message: fatalErr.Error()}) return } if err := cmd.Wait(); err != nil { _ = out.WriteJSON(CtrlMsg{Type: msgError, Message: fmt.Sprintf("sink process: %v (stderr: %s)", err, stderrBuf.String())}) return } _ = out.WriteJSON(CtrlMsg{Type: msgPushOK}) } func waitReady(fr *FrameReader, timeout time.Duration) error { type result struct { typ frameType err error } ch := make(chan result, 1) go func() { typ, _, err := fr.ReadFrame() ch <- result{typ, err} }() select { case r := <-ch: if r.err != nil { return fmt.Errorf("handshake failed: %w", r.err) } if r.typ != frameReady { return fmt.Errorf("handshake failed: unexpected frame type %d", r.typ) } return nil case <-time.After(timeout): return fmt.Errorf("handshake timed out after %s", timeout) } } type ackEvent struct { index uint64 eof bool err error } // pumpPush runs the source-side read/hash/compare/send loop against fw // (the pipe to the remote sink) while a single reader goroutine over fr // demultiplexes the three things the sink sends back on that one stream: // the streamed destination block hashes (fed to the source loop as they // arrive, so its own scan overlaps the sink's), the ACK/ERR frames for // blocks it wrote, and any relayed scan-progress. It returns once every // sent block's write has been confirmed. func pumpPush(req CtrlMsg, align int64, srcFile *os.File, fw *FrameWriter, fr *FrameReader, out *FrameWriter) error { blockCount := (req.Size + req.BlockSize - 1) / req.BlockSize const maxInFlight = 32 sem := make(chan struct{}, maxInFlight) var pendingMu sync.Mutex pending := make(map[uint64]bool) var copied, skipped int64 var lastProgress time.Time maybeProgress := func() { if time.Since(lastProgress) < 500*time.Millisecond { return } lastProgress = time.Now() _ = out.WriteJSON(CtrlMsg{ Type: msgProgress, Copied: atomic.LoadInt64(&copied), Skipped: atomic.LoadInt64(&skipped), TotalBlocks: blockCount, }) } // hashCh is sized to hold every block hash so the reader below never // blocks handing hashes off (which, since ACKs share the same stream, // would otherwise be able to deadlock against the in-flight-send limit). // This is the same order of memory the old whole-table transfer used. hashBuf := blockCount if hashBuf < 1 { hashBuf = 1 } hashCh := make(chan blockHash, hashBuf) ackEvents := make(chan ackEvent, 256) go func() { hashClosed := false closeHash := func() { if !hashClosed { close(hashCh) hashClosed = true } } for { typ, payload, err := fr.ReadFrame() if err != nil { closeHash() if err == io.EOF { ackEvents <- ackEvent{eof: true} } else { ackEvents <- ackEvent{err: err} } return } switch typ { case frameBlockHash: idx, h, derr := decodeBlockHashFrame(payload) if derr != nil { closeHash() ackEvents <- ackEvent{err: derr} return } hashCh <- blockHash{index: idx, hash: h} case frameHashDone: closeHash() case frameAck: idx, derr := decodeIndexFrame(payload) if derr != nil { closeHash() ackEvents <- ackEvent{err: derr} return } ackEvents <- ackEvent{index: idx} case frameErr: idx, msg, _ := decodeErrFrame(payload) closeHash() ackEvents <- ackEvent{err: fmt.Errorf("remote reported error at block %d: %s", idx, msg)} return case frameCtrlJSON: var m CtrlMsg if json.Unmarshal(payload, &m) == nil && m.Type == msgProgress { _ = out.WriteJSON(m) } default: closeHash() ackEvents <- ackEvent{err: fmt.Errorf("unexpected frame type %d from sink", typ)} return } } }() sendErrCh := make(chan error, 1) go func() { sendErrCh <- runSourceLoop(sourceLoopParams{ File: srcFile, Base: req.Base, Size: req.Size, BlockSize: req.BlockSize, Align: align, Hashes: hashCh, Out: fw, OnSkip: func(uint64) { atomic.AddInt64(&skipped, 1); maybeProgress() }, OnSend: func(idx uint64, _ [32]byte) { sem <- struct{}{} pendingMu.Lock() pending[idx] = true pendingMu.Unlock() }, }) }() // Wait for both: the send loop to finish (all DATA frames + DONE sent) // and the ack stream to end. Sink closes its stdout (a clean EOF) only // after it has acked every block it received, so an EOF while blocks // are still unconfirmed is treated as a real failure below. sendCh := sendErrCh ackCh := ackEvents var fatalErr error idle := time.NewTimer(120 * time.Second) defer idle.Stop() for sendCh != nil || ackCh != nil { if !idle.Stop() { select { case <-idle.C: default: } } idle.Reset(120 * time.Second) select { case sendErr := <-sendCh: sendCh = nil if sendErr != nil { fatalErr = sendErr } case ev := <-ackCh: switch { case ev.err != nil: fatalErr = ev.err ackCh = nil case ev.eof: ackCh = nil default: pendingMu.Lock() _, ok := pending[ev.index] delete(pending, ev.index) pendingMu.Unlock() if ok { atomic.AddInt64(&copied, 1) } <-sem maybeProgress() } case <-idle.C: fatalErr = fmt.Errorf("timed out waiting for the sink") } if fatalErr != nil { break } } if fatalErr != nil { return fatalErr } pendingMu.Lock() n := len(pending) pendingMu.Unlock() if n > 0 { return fmt.Errorf("sink closed the connection with %d block write confirmation(s) still outstanding", n) } _ = out.WriteJSON(CtrlMsg{Type: msgProgress, Copied: atomic.LoadInt64(&copied), Skipped: atomic.LoadInt64(&skipped), TotalBlocks: blockCount}) return nil } // --------------------------------------------------------------------- // pull driver: runs inside the DEST control agent. It reads and hashes the // local destination itself, spawns ssh to the source host running the // "source-stream" role, feeds it that hash table, then writes whatever it // streams back straight to the local destination — no round trip needed to // confirm a write, since dest-agent itself performed it. // --------------------------------------------------------------------- func runPullDriver(req CtrlMsg, out *FrameWriter) { tailArgs := []string{ "agent", "--role", "source-stream", "--path", req.PeerPath, "--base", strconv.FormatInt(req.PeerBase, 10), "--size", strconv.FormatInt(req.Size, 10), "--block-size", strconv.FormatInt(req.BlockSize, 10), } cmd := peerAgentCommand(req, tailArgs) stdin, err := cmd.StdinPipe() if err != nil { _ = out.WriteJSON(CtrlMsg{Type: msgPullFailed, Reason: err.Error()}) return } stdout, err := cmd.StdoutPipe() if err != nil { _ = out.WriteJSON(CtrlMsg{Type: msgPullFailed, Reason: err.Error()}) return } stderrBuf := newLimitedBuffer(4096) cmd.Stderr = stderrBuf if err := cmd.Start(); err != nil { _ = out.WriteJSON(CtrlMsg{Type: msgPullFailed, Reason: fmt.Sprintf("start ssh: %v", err)}) return } fw := NewFrameWriter(stdin) fr := NewFrameReader(stdout) timeout := time.Duration(req.ConnectTimeoutSec+2) * time.Second if err := waitReady(fr, timeout); err != nil { _ = cmd.Process.Kill() _ = cmd.Wait() _ = out.WriteJSON(CtrlMsg{Type: msgPullFailed, Reason: fmt.Sprintf("%v (remote stderr: %s)", err, stderrBuf.String())}) return } // Handshake succeeded. Open the destination and fingerprint its current // content block by block, streaming each hash to the source stream the // moment it is computed so it can start comparing straight away; write // whatever it streams back into the same file concurrently. dstFile, err := os.OpenFile(req.Path, os.O_RDWR, 0) if err != nil { _ = cmd.Process.Kill() _ = cmd.Wait() _ = out.WriteJSON(CtrlMsg{Type: msgError, Message: err.Error(), NeedPriv: isPermErr(err)}) return } defer dstFile.Close() align := alignmentFor(req.Path) hashErrCh := make(chan error, 1) go func() { err := streamHashBlocks(dstFile, req.Base, req.Size, req.BlockSize, align, func(bh blockHash) error { return fw.WriteFrame(frameBlockHash, encodeBlockHashFrame(bh.index, bh.hash)) }, scanProgressEmitter(out)) if err == nil { err = fw.WriteFrame(frameHashDone, nil) } if err != nil { // Unblock the source stream (waiting for more hashes) so the // dest loop below can unwind instead of hanging. _ = stdin.Close() } hashErrCh <- err }() loopErr := runDestLoop(destLoopParams{ File: dstFile, Base: req.Base, BlockSize: req.BlockSize, Align: align, In: fr, OnCtrlMsg: func(m CtrlMsg) { if m.Type == msgProgress { _ = out.WriteJSON(m) } }, }) hashErr := <-hashErrCh if hashErr != nil { _ = cmd.Process.Kill() _ = cmd.Wait() _ = out.WriteJSON(CtrlMsg{Type: msgError, Message: fmt.Sprintf("scan destination: %v", hashErr), NeedPriv: isPermErr(hashErr)}) return } if loopErr != nil { _ = cmd.Process.Kill() _ = cmd.Wait() _ = out.WriteJSON(CtrlMsg{Type: msgError, Message: loopErr.Error()}) return } if err := cmd.Wait(); err != nil { _ = out.WriteJSON(CtrlMsg{Type: msgError, Message: fmt.Sprintf("source-stream process: %v (stderr: %s)", err, stderrBuf.String())}) return } _ = out.WriteJSON(CtrlMsg{Type: msgPullOK}) } // --------------------------------------------------------------------- // sink role: one-shot process spawned (over ssh, in push mode) on the // destination host. Dumb write endpoint: verify+pwrite+ack per block. // --------------------------------------------------------------------- func runSinkRole(path string, base, size, blockSize int64) error { out := NewFrameWriter(os.Stdout) in := NewFrameReader(os.Stdin) f, err := os.OpenFile(path, os.O_RDWR, 0) if err != nil { fmt.Fprintf(os.Stderr, "sink: open %s: %v\n", path, err) return err } defer f.Close() align := alignmentFor(path) // Answer the handshake immediately so the push driver's short readiness // timeout isn't spent hashing a large destination. if err := out.WriteFrame(frameReady, nil); err != nil { return err } // Stream the destination's per-block hashes as they're computed while the // dest loop below concurrently receives and writes changed blocks. A // block is always hashed before a write for it can arrive (the source // only sends after seeing that block's hash), so the two accesses to f // never race on the same region. hashErrCh := make(chan error, 1) go func() { err := streamHashBlocks(f, base, size, blockSize, align, func(bh blockHash) error { return out.WriteFrame(frameBlockHash, encodeBlockHashFrame(bh.index, bh.hash)) }, scanProgressEmitter(out)) if err == nil { err = out.WriteFrame(frameHashDone, nil) } else { fmt.Fprintf(os.Stderr, "sink: hash %s: %v\n", path, err) _ = out.WriteFrame(frameErr, encodeErrFrame(0, err.Error())) } hashErrCh <- err }() loopErr := runDestLoop(destLoopParams{File: f, Base: base, BlockSize: blockSize, Align: align, In: in, AckOut: out}) if hashErr := <-hashErrCh; hashErr != nil { return hashErr } return loopErr } // --------------------------------------------------------------------- // source-stream role: one-shot process spawned (over ssh, in pull mode) on // the source host. Consumes the streamed destination block hashes, then // performs the same read/hash/compare/send loop a local push driver would, // writing straight to its own stdout. // --------------------------------------------------------------------- func runSourceStreamRole(path string, base, size, blockSize int64) error { out := NewFrameWriter(os.Stdout) in := NewFrameReader(os.Stdin) f, err := os.Open(path) if err != nil { fmt.Fprintf(os.Stderr, "source-stream: open %s: %v\n", path, err) return err } defer f.Close() align := alignmentFor(path) if err := out.WriteFrame(frameReady, nil); err != nil { return err } blockCount := (size + blockSize - 1) / blockSize hashBuf := blockCount if hashBuf < 1 { hashBuf = 1 } hashCh := make(chan blockHash, hashBuf) readErrCh := make(chan error, 1) go func() { for { typ, payload, err := in.ReadFrame() if err != nil { readErrCh <- fmt.Errorf("source-stream: read hash stream: %w", err) close(hashCh) return } switch typ { case frameBlockHash: idx, h, derr := decodeBlockHashFrame(payload) if derr != nil { readErrCh <- derr close(hashCh) return } hashCh <- blockHash{index: idx, hash: h} case frameHashDone: readErrCh <- nil close(hashCh) return default: readErrCh <- fmt.Errorf("source-stream: unexpected frame type %d", typ) close(hashCh) return } } }() var copied, skipped int64 var lastProgress time.Time maybeProgress := func() { if time.Since(lastProgress) < 500*time.Millisecond { return } lastProgress = time.Now() _ = out.WriteJSON(CtrlMsg{Type: msgProgress, Copied: copied, Skipped: skipped, TotalBlocks: blockCount}) } loopErr := runSourceLoop(sourceLoopParams{ File: f, Base: base, Size: size, BlockSize: blockSize, Align: align, Hashes: hashCh, Out: out, OnSkip: func(uint64) { skipped++; maybeProgress() }, OnSend: func(uint64, [32]byte) { copied++; maybeProgress() }, }) if readErr := <-readErrCh; readErr != nil { return readErr } return loopErr }