// Copyright (c) 2026 Petr BalvĂ­n (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") } }