417 lines
10 KiB
Go
417 lines
10 KiB
Go
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
|
|||
|
|
// SPDX-License-Identifier: MIT
|
||
|
|
|
||
|
|
package spmd
|
||
|
|
|
||
|
|
import (
|
||
|
|
"encoding/binary"
|
||
|
|
"errors"
|
||
|
|
"io"
|
||
|
|
"net"
|
||
|
|
"strings"
|
||
|
|
"sync"
|
||
|
|
"testing"
|
||
|
|
"time"
|
||
|
|
|
||
|
|
"sourcedock.dev/petrbalvin/tensor/internal/base"
|
||
|
|
)
|
||
|
|
|
||
|
|
// The world tests cover both transports with the same battery: an
|
||
|
|
// in-process world over channels and a loopback TCP world over real
|
||
|
|
// connections, because the contract says they are the same machine.
|
||
|
|
|
||
|
|
func TestLaunchOneRank(t *testing.T) {
|
||
|
|
err := Launch(1, func(w *World) error {
|
||
|
|
if w.Rank() != 0 || w.Size() != 1 {
|
||
|
|
t.Fatalf("rank %d of %d", w.Rank(), w.Size())
|
||
|
|
}
|
||
|
|
return w.Barrier()
|
||
|
|
})
|
||
|
|
if err != nil {
|
||
|
|
t.Fatal(err)
|
||
|
|
}
|
||
|
|
}
|
||
|
|
|
||
|
|
func TestLaunchBarrier(t *testing.T) {
|
||
|
|
for _, size := range []int{2, 3, 5, 8} {
|
||
|
|
t.Run("", func(t *testing.T) {
|
||
|
|
reached := make([]int, size)
|
||
|
|
err := Launch(size, func(w *World) error {
|
||
|
|
if err := w.Barrier(); err != nil {
|
||
|
|
return err
|
||
|
|
}
|
||
|
|
reached[w.Rank()] = 1
|
||
|
|
return w.Barrier()
|
||
|
|
})
|
||
|
|
if err != nil {
|
||
|
|
t.Fatal(err)
|
||
|
|
}
|
||
|
|
for r, got := range reached {
|
||
|
|
if got != 1 {
|
||
|
|
t.Fatalf("rank %d never reported", r)
|
||
|
|
}
|
||
|
|
}
|
||
|
|
})
|
||
|
|
}
|
||
|
|
}
|
||
|
|
|
||
|
|
// TestLaunchFailsTogether is the fail-fast rule: the rank that sees the
|
||
|
|
// error fails its world, every other rank's next collective answers
|
||
|
|
// with an error, and Launch returns the first rank's error in rank
|
||
|
|
// order.
|
||
|
|
func TestLaunchFailsTogether(t *testing.T) {
|
||
|
|
boom := errors.New("boom")
|
||
|
|
saw := make([]error, 3)
|
||
|
|
err := Launch(3, func(w *World) error {
|
||
|
|
if w.Rank() == 1 {
|
||
|
|
return boom
|
||
|
|
}
|
||
|
|
saw[w.Rank()] = w.Barrier()
|
||
|
|
return saw[w.Rank()]
|
||
|
|
})
|
||
|
|
if err == nil {
|
||
|
|
t.Fatal("a world with a failing rank returned nil")
|
||
|
|
}
|
||
|
|
for r, got := range saw {
|
||
|
|
if r == 1 {
|
||
|
|
continue
|
||
|
|
}
|
||
|
|
if got == nil {
|
||
|
|
t.Fatalf("rank %d's barrier survived a failed world", r)
|
||
|
|
}
|
||
|
|
}
|
||
|
|
}
|
||
|
|
|
||
|
|
func TestLaunchPanicIsAnError(t *testing.T) {
|
||
|
|
err := Launch(2, func(w *World) error {
|
||
|
|
if w.Rank() == 1 {
|
||
|
|
panic("rank one fell over")
|
||
|
|
}
|
||
|
|
return w.Barrier()
|
||
|
|
})
|
||
|
|
if err == nil || !strings.Contains(err.Error(), "rank 1 panicked") {
|
||
|
|
t.Fatalf("a panicking rank came back as %v", err)
|
||
|
|
}
|
||
|
|
}
|
||
|
|
|
||
|
|
// runTCPWorld assembles a loopback TCP world of size ranks, running fn
|
||
|
|
// on every rank in its own goroutine, and fails the test if any rank
|
||
|
|
// errors. The listener the tests hand over lets rank 0 know its
|
||
|
|
// address before it starts. A final barrier runs after fn on every
|
||
|
|
// rank, the orderly end of the program: no rank tears its world down
|
||
|
|
// while another still expects words from it.
|
||
|
|
func runTCPWorld(t *testing.T, size int, opts Options, fn func(w *World) error) {
|
||
|
|
t.Helper()
|
||
|
|
ln, err := net.Listen("tcp", "127.0.0.1:0")
|
||
|
|
if err != nil {
|
||
|
|
t.Fatal(err)
|
||
|
|
}
|
||
|
|
defer ln.Close()
|
||
|
|
var wg sync.WaitGroup
|
||
|
|
errs := make([]error, size)
|
||
|
|
wg.Go(func() {
|
||
|
|
w, err := listen(ln, size, opts)
|
||
|
|
if err != nil {
|
||
|
|
errs[0] = err
|
||
|
|
return
|
||
|
|
}
|
||
|
|
defer w.Close()
|
||
|
|
if err := fn(w); err != nil {
|
||
|
|
errs[0] = err
|
||
|
|
return
|
||
|
|
}
|
||
|
|
errs[0] = w.Barrier()
|
||
|
|
})
|
||
|
|
for r := 1; r < size; r++ {
|
||
|
|
wg.Go(func() {
|
||
|
|
w, err := Join(ln.Addr().String(), opts)
|
||
|
|
if err != nil {
|
||
|
|
errs[r] = err
|
||
|
|
return
|
||
|
|
}
|
||
|
|
defer w.Close()
|
||
|
|
if err := fn(w); err != nil {
|
||
|
|
errs[r] = err
|
||
|
|
return
|
||
|
|
}
|
||
|
|
errs[r] = w.Barrier()
|
||
|
|
})
|
||
|
|
}
|
||
|
|
wg.Wait()
|
||
|
|
for r, err := range errs {
|
||
|
|
if err != nil {
|
||
|
|
t.Fatalf("rank %d: %v", r, err)
|
||
|
|
}
|
||
|
|
}
|
||
|
|
}
|
||
|
|
|
||
|
|
func TestTCPBarrier(t *testing.T) {
|
||
|
|
runTCPWorld(t, 4, Options{Timeout: 30 * time.Second}, func(w *World) error {
|
||
|
|
if err := w.Barrier(); err != nil {
|
||
|
|
return err
|
||
|
|
}
|
||
|
|
return w.Barrier()
|
||
|
|
})
|
||
|
|
}
|
||
|
|
|
||
|
|
// TestTCPRankOrderIsDialOrder pins the one place ranks come from: the
|
||
|
|
// order the peers dial in, which the collectives' results never depend
|
||
|
|
// on.
|
||
|
|
func TestTCPRankOrderIsDialOrder(t *testing.T) {
|
||
|
|
runTCPWorld(t, 3, Options{Timeout: 30 * time.Second}, func(w *World) error {
|
||
|
|
return w.Barrier()
|
||
|
|
})
|
||
|
|
}
|
||
|
|
|
||
|
|
// TestTCPStrayConnectionTakesNoRank dials the listener with a
|
||
|
|
// connection that says nothing the handshake would recognise; the world
|
||
|
|
// must still assemble on the real ranks.
|
||
|
|
func TestTCPStrayConnectionTakesNoRank(t *testing.T) {
|
||
|
|
ln, err := net.Listen("tcp", "127.0.0.1:0")
|
||
|
|
if err != nil {
|
||
|
|
t.Fatal(err)
|
||
|
|
}
|
||
|
|
defer ln.Close()
|
||
|
|
// The stray speaks first, then goes silent.
|
||
|
|
stray, err := net.Dial("tcp", ln.Addr().String())
|
||
|
|
if err != nil {
|
||
|
|
t.Fatal(err)
|
||
|
|
}
|
||
|
|
defer stray.Close()
|
||
|
|
if _, err := stray.Write([]byte("not the magic at all, but long enough to fill the read")); err != nil {
|
||
|
|
t.Fatal(err)
|
||
|
|
}
|
||
|
|
done := make(chan struct{})
|
||
|
|
var worldErr error
|
||
|
|
go func() {
|
||
|
|
defer close(done)
|
||
|
|
w, err := listen(ln, 2, Options{Timeout: 30 * time.Second})
|
||
|
|
if err != nil {
|
||
|
|
worldErr = err
|
||
|
|
return
|
||
|
|
}
|
||
|
|
w.Close()
|
||
|
|
}()
|
||
|
|
// The real rank joins behind the stray.
|
||
|
|
joined := make(chan error, 1)
|
||
|
|
go func() {
|
||
|
|
w, err := Join(ln.Addr().String(), Options{Timeout: 30 * time.Second})
|
||
|
|
if err != nil {
|
||
|
|
joined <- err
|
||
|
|
return
|
||
|
|
}
|
||
|
|
w.Close()
|
||
|
|
joined <- nil
|
||
|
|
}()
|
||
|
|
select {
|
||
|
|
case err := <-joined:
|
||
|
|
if err != nil {
|
||
|
|
t.Fatalf("the real rank did not join behind the stray: %v", err)
|
||
|
|
}
|
||
|
|
case <-time.After(30 * time.Second):
|
||
|
|
t.Fatal("the world never assembled")
|
||
|
|
}
|
||
|
|
<-done
|
||
|
|
if worldErr != nil {
|
||
|
|
t.Fatalf("rank 0: %v", worldErr)
|
||
|
|
}
|
||
|
|
}
|
||
|
|
|
||
|
|
// TestTCPDeadlineFailsTheCollective is the stuck-peer rule: a rank
|
||
|
|
// whose peer stops answering is errored out by its own deadline, never
|
||
|
|
// left hanging. The peer's world carries the short timeout from birth,
|
||
|
|
// so nothing changes under a running pump.
|
||
|
|
func TestTCPDeadlineFailsTheCollective(t *testing.T) {
|
||
|
|
ln, err := net.Listen("tcp", "127.0.0.1:0")
|
||
|
|
if err != nil {
|
||
|
|
t.Fatal(err)
|
||
|
|
}
|
||
|
|
defer ln.Close()
|
||
|
|
hubDone := make(chan *World, 1)
|
||
|
|
go func() {
|
||
|
|
// The hub assembles and stays silent: it never answers.
|
||
|
|
w, err := listen(ln, 2, Options{Timeout: 30 * time.Second})
|
||
|
|
if err != nil {
|
||
|
|
w = nil
|
||
|
|
}
|
||
|
|
hubDone <- w
|
||
|
|
}()
|
||
|
|
peer, err := Join(ln.Addr().String(), Options{Timeout: 80 * time.Millisecond})
|
||
|
|
if err != nil {
|
||
|
|
t.Fatal(err)
|
||
|
|
}
|
||
|
|
defer peer.Close()
|
||
|
|
hub := <-hubDone
|
||
|
|
if hub == nil {
|
||
|
|
t.Fatal("the hub did not assemble")
|
||
|
|
}
|
||
|
|
defer hub.Close()
|
||
|
|
if err := peer.Barrier(); err == nil {
|
||
|
|
t.Fatal("a barrier against a silent hub succeeded")
|
||
|
|
}
|
||
|
|
}
|
||
|
|
|
||
|
|
// TestTCPMaxMessageRefused: a peer announcing a payload beyond the
|
||
|
|
// ceiling is an error before any allocation.
|
||
|
|
func TestTCPMaxMessageRefused(t *testing.T) {
|
||
|
|
ln, err := net.Listen("tcp", "127.0.0.1:0")
|
||
|
|
if err != nil {
|
||
|
|
t.Fatal(err)
|
||
|
|
}
|
||
|
|
defer ln.Close()
|
||
|
|
const ceiling = int64(1 << 20)
|
||
|
|
wCh := make(chan *World, 1)
|
||
|
|
go func() {
|
||
|
|
w, err := listen(ln, 2, Options{Timeout: 30 * time.Second, MaxMessage: ceiling})
|
||
|
|
if err != nil {
|
||
|
|
wCh <- nil
|
||
|
|
return
|
||
|
|
}
|
||
|
|
wCh <- w
|
||
|
|
}()
|
||
|
|
// The fake peer does the handshake by hand, then announces an
|
||
|
|
// oversized frame.
|
||
|
|
peerDone := make(chan error, 1)
|
||
|
|
go func() {
|
||
|
|
conn, err := net.Dial("tcp", ln.Addr().String())
|
||
|
|
if err != nil {
|
||
|
|
peerDone <- err
|
||
|
|
return
|
||
|
|
}
|
||
|
|
defer conn.Close()
|
||
|
|
if err := sendHello(conn); err != nil {
|
||
|
|
peerDone <- err
|
||
|
|
return
|
||
|
|
}
|
||
|
|
if _, _, err := readWelcome(conn, time.Now().Add(30*time.Second)); err != nil {
|
||
|
|
peerDone <- err
|
||
|
|
return
|
||
|
|
}
|
||
|
|
link := newTCPLink(conn)
|
||
|
|
var head [frameHeaderLen]byte
|
||
|
|
head[0] = tagBarrier
|
||
|
|
head[1] = frameVersion
|
||
|
|
binary.LittleEndian.PutUint32(head[4:], 1) // from
|
||
|
|
binary.LittleEndian.PutUint32(head[8:], 0) // dest
|
||
|
|
binary.LittleEndian.PutUint64(head[12:], uint64(ceiling+1))
|
||
|
|
if _, err := link.wr.Write(head[:]); err != nil {
|
||
|
|
peerDone <- err
|
||
|
|
return
|
||
|
|
}
|
||
|
|
peerDone <- link.wr.Flush()
|
||
|
|
}()
|
||
|
|
w0 := <-wCh
|
||
|
|
if w0 == nil {
|
||
|
|
t.Fatal("rank 0 did not assemble")
|
||
|
|
}
|
||
|
|
defer w0.Close()
|
||
|
|
if err := <-peerDone; err != nil {
|
||
|
|
t.Fatalf("the fake peer: %v", err)
|
||
|
|
}
|
||
|
|
if _, err := w0.recvFrom(1, tagBarrier); err == nil {
|
||
|
|
t.Fatal("an oversized frame was received")
|
||
|
|
}
|
||
|
|
}
|
||
|
|
|
||
|
|
// TestHandshake walks the joining words over an in-memory connection:
|
||
|
|
// the right magic passes both ways, a wrong magic is refused, and an
|
||
|
|
// impossible rank answer is refused.
|
||
|
|
func TestHandshake(t *testing.T) {
|
||
|
|
c, s := net.Pipe()
|
||
|
|
defer c.Close()
|
||
|
|
defer s.Close()
|
||
|
|
deadline := time.Now().Add(5 * time.Second)
|
||
|
|
go sendHello(c)
|
||
|
|
if err := readHello(s, deadline); err != nil {
|
||
|
|
t.Fatalf("the right magic was refused: %v", err)
|
||
|
|
}
|
||
|
|
go sendWelcome(c, 4, 2)
|
||
|
|
size, rank, err := readWelcome(s, deadline)
|
||
|
|
if err != nil {
|
||
|
|
t.Fatalf("the welcome did not read: %v", err)
|
||
|
|
}
|
||
|
|
if size != 4 || rank != 2 {
|
||
|
|
t.Fatalf("welcome answered size %d rank %d", size, rank)
|
||
|
|
}
|
||
|
|
var bad [handshakeLen]byte
|
||
|
|
copy(bad[0:4], []byte("XXXX"))
|
||
|
|
bad[4] = frameVersion
|
||
|
|
go c.Write(bad[:])
|
||
|
|
if err := readHello(s, deadline); err == nil {
|
||
|
|
t.Fatal("a wrong magic was accepted")
|
||
|
|
}
|
||
|
|
go sendWelcome(c, 4, 4)
|
||
|
|
if _, _, err := readWelcome(s, deadline); err == nil {
|
||
|
|
t.Fatal("a rank beyond the world's size was accepted")
|
||
|
|
}
|
||
|
|
}
|
||
|
|
|
||
|
|
// The compile-time guards on the error surface: every failure this
|
||
|
|
// package reports keeps the library's prefix.
|
||
|
|
func TestErrorPrefix(t *testing.T) {
|
||
|
|
err := base.Errf("spmd: test")
|
||
|
|
if err == nil || !strings.HasPrefix(err.Error(), "tensor: ") {
|
||
|
|
t.Fatalf("the package error lost its prefix: %v", err)
|
||
|
|
}
|
||
|
|
var _ io.Closer = (*World)(nil)
|
||
|
|
}
|
||
|
|
|
||
|
|
// TestTCPNegativeLengthRefused: a frame length with its top bit set
|
||
|
|
// turns negative through the signed conversion; the receiver refuses
|
||
|
|
// it instead of allocating from it, which once panicked the hub.
|
||
|
|
func TestTCPNegativeLengthRefused(t *testing.T) {
|
||
|
|
ln, err := net.Listen("tcp", "127.0.0.1:0")
|
||
|
|
if err != nil {
|
||
|
|
t.Fatal(err)
|
||
|
|
}
|
||
|
|
defer ln.Close()
|
||
|
|
wCh := make(chan *World, 1)
|
||
|
|
go func() {
|
||
|
|
w, err := listen(ln, 2, Options{Timeout: 30 * time.Second})
|
||
|
|
if err != nil {
|
||
|
|
w = nil
|
||
|
|
}
|
||
|
|
wCh <- w
|
||
|
|
}()
|
||
|
|
peerDone := make(chan error, 1)
|
||
|
|
go func() {
|
||
|
|
conn, err := net.Dial("tcp", ln.Addr().String())
|
||
|
|
if err != nil {
|
||
|
|
peerDone <- err
|
||
|
|
return
|
||
|
|
}
|
||
|
|
defer conn.Close()
|
||
|
|
if err := sendHello(conn); err != nil {
|
||
|
|
peerDone <- err
|
||
|
|
return
|
||
|
|
}
|
||
|
|
if _, _, err := readWelcome(conn, time.Now().Add(30*time.Second)); err != nil {
|
||
|
|
peerDone <- err
|
||
|
|
return
|
||
|
|
}
|
||
|
|
link := newTCPLink(conn)
|
||
|
|
var head [frameHeaderLen]byte
|
||
|
|
head[0] = tagBarrier
|
||
|
|
head[1] = frameVersion
|
||
|
|
binary.LittleEndian.PutUint32(head[4:], 1)
|
||
|
|
binary.LittleEndian.PutUint32(head[8:], 0)
|
||
|
|
binary.LittleEndian.PutUint64(head[12:], uint64(1)<<63)
|
||
|
|
if _, err := link.wr.Write(head[:]); err != nil {
|
||
|
|
peerDone <- err
|
||
|
|
return
|
||
|
|
}
|
||
|
|
peerDone <- link.wr.Flush()
|
||
|
|
}()
|
||
|
|
w0 := <-wCh
|
||
|
|
if w0 == nil {
|
||
|
|
t.Fatal("rank 0 did not assemble")
|
||
|
|
}
|
||
|
|
defer w0.Close()
|
||
|
|
if err := <-peerDone; err != nil {
|
||
|
|
t.Fatalf("the fake peer: %v", err)
|
||
|
|
}
|
||
|
|
if _, err := w0.recvFrom(1, tagBarrier); err == nil {
|
||
|
|
t.Fatal("a negative length was received")
|
||
|
|
}
|
||
|
|
}
|