package main import ( "encoding/hex" "encoding/json" "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 (internal)") path := fs.String("path", "", "path to read/write") size := fs.Int64("size", 0, "total sync size in bytes") blockSize := fs.Int64("block-size", defaultBlockSize, "block size in bytes") if err := fs.Parse(args); err != nil { return err } switch *role { case "control": return runControlAgent() case "sink": return runSinkRole(*path, *size, *blockSize) case "source-stream": return runSourceStreamRole(*path, *size, *blockSize) default: return fmt.Errorf("agent: unknown or missing --role %q (want control|sink|source-stream)", *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: hashes, err := readHashTable(in) if err != nil { return err } runPushDriver(m, hashes, out) case msgConnectPull: hashes, err := readHashTable(in) if err != nil { return err } runPullDriver(m, hashes, out) case msgClose: _ = out.WriteJSON(CtrlMsg{Type: msgBye}) return nil default: _ = out.WriteJSON(CtrlMsg{Type: msgError, Message: fmt.Sprintf("unknown command %q", m.Type)}) } } } func readHashTable(in *FrameReader) ([][32]byte, error) { typ, payload, err := in.ReadFrame() if err != nil { return nil, fmt.Errorf("read hash table: %w", err) } if typ != frameHashTable { return nil, fmt.Errorf("expected hash table frame, got type %d", typ) } return unflattenHashes(payload) } func handleStat(out *FrameWriter, m CtrlMsg) { info, err := statPath(m.Path) if err != nil { _ = out.WriteJSON(CtrlMsg{Type: msgError, Message: err.Error()}) 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()}) return } _ = out.WriteJSON(CtrlMsg{Type: msgPrepareOK, Size: m.Size}) } // --------------------------------------------------------------------- // 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, hashes [][32]byte, out *FrameWriter) { tailArgs := []string{ "agent", "--role", "sink", "--path", req.PeerPath, "--size", strconv.FormatInt(req.Size, 10), "--block-size", strconv.FormatInt(req.BlockSize, 10), } var cmd *exec.Cmd if req.PeerLocal { cmd = localAgentCommand(tailArgs) } else { cmd = sshCommand(req.SSHBin, req.SSHOpts, true, req.ConnectTimeoutSec, req.PeerUser, req.PeerHost, append([]string{req.RemoteBin}, 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. srcFile, err := os.Open(req.Path) if err != nil { _ = cmd.Process.Kill() _ = cmd.Wait() _ = out.WriteJSON(CtrlMsg{Type: msgError, Message: err.Error()}) return } defer srcFile.Close() if fatalErr := pumpPush(req, hashes, 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 concurrently draining ACK/ERR frames // from fr, only reporting a block as done to the manager (over out) once // its write has been confirmed. func pumpPush(req CtrlMsg, hashes [][32]byte, srcFile *os.File, fw *FrameWriter, fr *FrameReader, out *FrameWriter) error { const maxInFlight = 32 sem := make(chan struct{}, maxInFlight) var pendingMu sync.Mutex pending := make(map[uint64][32]byte) batch := newResultBatcher(out) defer batch.flush() 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: int64(len(hashes)), }) } ackEvents := make(chan ackEvent, 256) go func() { for { typ, payload, err := fr.ReadFrame() if err != nil { if err == io.EOF { ackEvents <- ackEvent{eof: true} } else { ackEvents <- ackEvent{err: err} } return } switch typ { case frameAck: idx, err := decodeIndexFrame(payload) if err != nil { ackEvents <- ackEvent{err: err} return } ackEvents <- ackEvent{index: idx} case frameErr: idx, msg, _ := decodeErrFrame(payload) ackEvents <- ackEvent{err: fmt.Errorf("remote reported error at block %d: %s", idx, msg)} return default: 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, Size: req.Size, BlockSize: req.BlockSize, Hashes: hashes, Out: fw, OnSkip: func(uint64) { atomic.AddInt64(&skipped, 1); maybeProgress() }, OnSend: func(idx uint64, hash [32]byte) { sem <- struct{}{} pendingMu.Lock() pending[idx] = hash 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() h, ok := pending[ev.index] delete(pending, ev.index) pendingMu.Unlock() if ok { atomic.AddInt64(&copied, 1) batch.add(BlockResult{Index: ev.index, Hash: hex.EncodeToString(h[:])}) } <-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) } batch.flush() _ = out.WriteJSON(CtrlMsg{Type: msgProgress, Copied: atomic.LoadInt64(&copied), Skipped: atomic.LoadInt64(&skipped), TotalBlocks: int64(len(hashes))}) return nil } // --------------------------------------------------------------------- // pull driver: runs inside the DEST control agent. Spawns ssh straight to // the source host, running the "source-stream" role, feeds it the hash // table, then writes whatever it streams back directly to the local // destination — no round trip needed to confirm a write, since dest-agent // itself performed it. // --------------------------------------------------------------------- func runPullDriver(req CtrlMsg, hashes [][32]byte, out *FrameWriter) { tailArgs := []string{ "agent", "--role", "source-stream", "--path", req.PeerPath, "--size", strconv.FormatInt(req.Size, 10), "--block-size", strconv.FormatInt(req.BlockSize, 10), } var cmd *exec.Cmd if req.PeerLocal { cmd = localAgentCommand(tailArgs) } else { cmd = sshCommand(req.SSHBin, req.SSHOpts, true, req.ConnectTimeoutSec, req.PeerUser, req.PeerHost, append([]string{req.RemoteBin}, 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: send the hash table and commit to pull. if err := fw.WriteFrame(frameHashTable, flattenHashes(hashes)); err != nil { _ = cmd.Process.Kill() _ = cmd.Wait() _ = out.WriteJSON(CtrlMsg{Type: msgError, Message: fmt.Sprintf("send hash table: %v", err)}) return } 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()}) return } defer dstFile.Close() batch := newResultBatcher(out) var copied int64 loopErr := runDestLoop(destLoopParams{ File: dstFile, BlockSize: req.BlockSize, In: fr, OnWritten: func(idx uint64, hash [32]byte) { atomic.AddInt64(&copied, 1) batch.add(BlockResult{Index: idx, Hash: hex.EncodeToString(hash[:])}) }, OnCtrlMsg: func(m CtrlMsg) { if m.Type == msgProgress { _ = out.WriteJSON(m) } }, }) batch.flush() 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, 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() if err := out.WriteFrame(frameReady, nil); err != nil { return err } return runDestLoop(destLoopParams{File: f, BlockSize: blockSize, In: in, AckOut: out}) } // --------------------------------------------------------------------- // source-stream role: one-shot process spawned (over ssh, in pull mode) on // the source host. Reads the hash table, then performs the same // read/hash/compare/send loop a local push driver would, writing straight // to its own stdout. // --------------------------------------------------------------------- func runSourceStreamRole(path string, 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() if err := out.WriteFrame(frameReady, nil); err != nil { return err } typ, payload, err := in.ReadFrame() if err != nil { return fmt.Errorf("source-stream: read hash table: %w", err) } if typ != frameHashTable { return fmt.Errorf("source-stream: expected hash table frame, got type %d", typ) } hashes, err := unflattenHashes(payload) if err != nil { return err } 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: int64(len(hashes))}) } return runSourceLoop(sourceLoopParams{ File: f, Size: size, BlockSize: blockSize, Hashes: hashes, Out: out, OnSkip: func(uint64) { skipped++; maybeProgress() }, OnSend: func(uint64, [32]byte) { copied++; maybeProgress() }, }) } // --------------------------------------------------------------------- // resultBatcher: coalesces BlockResult entries into occasional // block_done_batch messages so the manager isn't hit with one control // message per synced block. // --------------------------------------------------------------------- type resultBatcher struct { out *FrameWriter mu sync.Mutex buf []BlockResult last time.Time } func newResultBatcher(out *FrameWriter) *resultBatcher { return &resultBatcher{out: out, last: time.Now()} } func (b *resultBatcher) add(r BlockResult) { b.mu.Lock() b.buf = append(b.buf, r) shouldFlush := len(b.buf) >= 256 || time.Since(b.last) > 500*time.Millisecond b.mu.Unlock() if shouldFlush { b.flush() } } func (b *resultBatcher) flush() { b.mu.Lock() if len(b.buf) == 0 { b.mu.Unlock() return } entries := b.buf b.buf = nil b.last = time.Now() b.mu.Unlock() _ = b.out.WriteJSON(CtrlMsg{Type: msgBlockDoneBatch, Entries: entries}) }