clonetool/agent.go
2026-09-05 23:02:29 +02:00

558 lines
16 KiB
Go

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})
}