Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
86 changes: 86 additions & 0 deletions internal/protocol/protocol_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -512,3 +512,89 @@ func TestCancelledContextStopsClient(t *testing.T) {
t.Error("Finish ignored a cancelled context")
}
}

func TestResumeOverPipes(t *testing.T) {
src := t.TempDir()
payload := randomBytes(400<<10, 6)
writeFile(t, src, "model.safetensors", payload)
m, err := scan.Build(context.Background(), scan.Options{Root: src, Tool: "test"})
if err != nil {
t.Fatal(err)
}

dst := filepath.Join(t.TempDir(), "out")
stage := filepath.Join(dst, manifest.StateDir, "stage", "model.safetensors.part")
if err := os.MkdirAll(filepath.Dir(stage), 0o755); err != nil {
t.Fatal(err)
}
partial := make([]byte, len(payload))
f := m.Files[0]
var filled int
for _, c := range f.Chunks {
if filled >= len(payload)/2 {
break
}
copy(partial[c.Offset:], payload[c.Offset:c.End()])
filled += int(c.Length)
}
if err := os.WriteFile(stage, partial, 0o644); err != nil {
t.Fatal(err)
}

client, done := pipePair(t, ServerOptions{Root: dst, Tool: "modelmove/test"})
plan, _ := runTransfer(t, client, m, src, defaultRequest())
done()

if plan.NeedBytes >= int64(len(payload)) {
t.Errorf("resume sent %d of %d bytes; staged chunks should have been kept", plan.NeedBytes, len(payload))
}
got, err := os.ReadFile(filepath.Join(dst, "model.safetensors"))
if err != nil {
t.Fatal(err)
}
if !bytes.Equal(got, payload) {
t.Fatal("the resumed file does not match the source")
}
}

func TestCorruptionRepairOverPipes(t *testing.T) {
src := sourceDir(t)
dst := filepath.Join(t.TempDir(), "out")
m, err := scan.Build(context.Background(), scan.Options{Root: src, Tool: "test"})
if err != nil {
t.Fatal(err)
}

client, done := pipePair(t, ServerOptions{Root: dst, Tool: "modelmove/test"})
runTransfer(t, client, m, src, defaultRequest())
done()

rel := "model-00002-of-00002.safetensors"
path := filepath.Join(dst, rel)
data, err := os.ReadFile(path)
if err != nil {
t.Fatal(err)
}
data[len(data)/2] ^= 0xff
if err := os.WriteFile(path, data, 0o644); err != nil {
t.Fatal(err)
}

client, done = pipePair(t, ServerOptions{Root: dst, Tool: "modelmove/test"})
plan, _ := runTransfer(t, client, m, src, defaultRequest())
done()
if plan.NeedBytes == 0 {
t.Fatal("repair planned no bytes after dest corruption")
}
want, err := os.ReadFile(filepath.Join(src, rel))
if err != nil {
t.Fatal(err)
}
got, err := os.ReadFile(path)
if err != nil {
t.Fatal(err)
}
if !bytes.Equal(want, got) {
t.Fatal("repair did not restore the source bytes")
}
}
48 changes: 47 additions & 1 deletion scripts/e2e-live.sh
Original file line number Diff line number Diff line change
Expand Up @@ -20,7 +20,7 @@ if ! "${SSH[@]}" "${REMOTE_USER}@${REMOTE_HOST}" true >/dev/null 2>&1; then
fi

WORK=$(mktemp -d)
trap 'rm -rf "$WORK"; "${SSH[@]}" "${REMOTE_USER}@${REMOTE_HOST}" "rm -rf /tmp/modelmove-live-$$" >/dev/null 2>&1 || true' EXIT
trap 'rm -rf "$WORK"; "${SSH[@]}" "${REMOTE_USER}@${REMOTE_HOST}" "rm -rf /tmp/modelmove-live-$$ /tmp/modelmove-live-resume-$$" >/dev/null 2>&1 || true' EXIT

REMOTE_DST="/tmp/modelmove-live-$$"
SRC="$WORK/src"
Expand Down Expand Up @@ -81,5 +81,51 @@ PY

"${SSH[@]}" "${REMOTE_USER}@${REMOTE_HOST}" "'$BIN' verify '$REMOTE_DST' --no-progress"

echo "==> e2e-live: corruption is detected and repaired over live sshd"
"${SSH[@]}" "${REMOTE_USER}@${REMOTE_HOST}" "python3 - '$REMOTE_DST/model-00002-of-00002.safetensors'" <<'PY'
import sys
p = sys.argv[1]
d = bytearray(open(p, "rb").read())
d[1_000_000:1_000_032] = b"C" * 32
open(p, "wb").write(bytes(d))
PY
set +e
"${SSH[@]}" "${REMOTE_USER}@${REMOTE_HOST}" "'$BIN' verify '$REMOTE_DST' --no-progress" > "$WORK/verify.out" 2>&1
code=$?
set -e
[ "$code" -eq 2 ] || fail "verify exited $code on a corrupt model, want 2"
grep -q "bad chunk" "$WORK/verify.out" || fail "verify did not locate the bad chunk"

"$BIN" sync "$SRC" "$TARGET" --remote-bin "$BIN" --no-progress
"${SSH[@]}" "${REMOTE_USER}@${REMOTE_HOST}" "'$BIN' verify '$REMOTE_DST' --no-progress"

echo "==> e2e-live: resume reuses a planted staging file"
RESUME_DST="/tmp/modelmove-live-resume-$$"
RESUME_TARGET="${REMOTE_USER}@${REMOTE_HOST}:${RESUME_DST}"
python3 - "$SRC/model-00001-of-00002.safetensors" "$WORK/partial.part" <<'PY'
import sys
src, dest = sys.argv[1], sys.argv[2]
data = open(src, "rb").read()
partial = bytearray(len(data))
# First 2 MiB is enough for the 4 MiB shard to show reuse without
# needing the FastCDC cut points on this side.
partial[:2_000_000] = data[:2_000_000]
open(dest, "wb").write(partial)
PY
"${SSH[@]}" "${REMOTE_USER}@${REMOTE_HOST}" "mkdir -p '$RESUME_DST/.modelmove/stage'"
scp -o BatchMode=yes -o ConnectTimeout=2 "$WORK/partial.part" \
"${REMOTE_USER}@${REMOTE_HOST}:${RESUME_DST}/.modelmove/stage/model-00001-of-00002.safetensors.part"
"$BIN" copy "$SRC" "$RESUME_TARGET" --remote-bin "$BIN" --json --no-progress > "$WORK/resume.json"
python3 - "$WORK/resume.json" <<'PY'
import json, sys
r = json.load(open(sys.argv[1]))
total = r["plan"]["total_bytes"]
need = r["plan"]["need_bytes"]
assert need < total, f"resume planned {need} of {total} bytes; staged prefix should have been reused"
print(f" live ssh resume: planned {need} of {total} bytes")
PY
"${SSH[@]}" "${REMOTE_USER}@${REMOTE_HOST}" "'$BIN' verify '$RESUME_DST' --no-progress"
"${SSH[@]}" "${REMOTE_USER}@${REMOTE_HOST}" "rm -rf '$RESUME_DST'"

echo
echo "e2e-live: all checks passed"