Files

417 lines
10 KiB
Go
Raw Permalink Normal View History

2026-09-03 10:00:00 +02:00
// 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")
}
}