// Copyright (c) 2026 Petr BalvĂ­n (https://petrbalvin.org) // SPDX-License-Identifier: MIT package spmd import ( "encoding/binary" "runtime" "testing" "time" "sourcedock.dev/petrbalvin/tensor/internal/base" ) // The pool tests cover the two bounds, the exact-length answer, the // switch, the bound's flatness in the process's heap, and the one // ownership chain itself: routed frames carrying distinct bits through // a real hub while the pool recycles every buffer between pump and // drain. // resetPool empties the shared pool so no test inherits another's // buffers, pins the switch on, and restores both at the end. func resetPool(t *testing.T) { t.Helper() was := framePoolEnabled t.Cleanup(func() { framePoolEnabled = was emptyPool() }) emptyPool() framePoolEnabled = true } func emptyPool() { routedFrames.mu.Lock() routedFrames.free = nil routedFrames.mu.Unlock() } // poolDepth answers how many buffers the pool holds. func poolDepth() int { routedFrames.mu.Lock() defer routedFrames.mu.Unlock() return len(routedFrames.free) } func TestFramePoolTakeAnswersExactLengths(t *testing.T) { resetPool(t) if buf := routedFrames.take(1000); buf != nil { t.Fatal("an empty pool answered a buffer") } routedFrames.retire(make([]byte, 1000)) buf := routedFrames.take(600) if buf == nil { t.Fatal("a retained buffer was not answered") } if len(buf) != 600 || cap(buf) < 600 { t.Fatalf("take answered len %d cap %d for a 600 byte payload", len(buf), cap(buf)) } if buf := routedFrames.take(600); buf != nil { t.Fatal("a taken buffer was answered twice") } } func TestFramePoolTakesTheSmallestThatFits(t *testing.T) { resetPool(t) routedFrames.retire(make([]byte, 5000)) routedFrames.retire(make([]byte, 900)) buf := routedFrames.take(800) if buf == nil { t.Fatal("a retained buffer was not answered") } if cap(buf) != 900 { t.Fatalf("take answered cap %d when a 900 byte buffer was retained", cap(buf)) } buf = routedFrames.take(800) if buf == nil || cap(buf) != 5000 { t.Fatalf("the second take answered cap %d, want the 5000 byte buffer", cap(buf)) } } func TestFramePoolRefusesTheOversized(t *testing.T) { resetPool(t) if buf := routedFrames.take(framePoolCeiling + 1); buf != nil { t.Fatal("a length beyond the ceiling was answered") } routedFrames.retire(make([]byte, framePoolCeiling+1)) if got := poolDepth(); got != 0 { t.Fatalf("a buffer beyond the ceiling was retained, pool holds %d", got) } routedFrames.retire(make([]byte, framePoolCeiling)) if got := poolDepth(); got != 1 { t.Fatalf("a buffer at the ceiling was refused, pool holds %d", got) } if buf := routedFrames.take(framePoolCeiling); buf == nil { t.Fatal("a length at the ceiling was not answered") } } func TestFramePoolHoldsTheBound(t *testing.T) { resetPool(t) for range framePoolBound + 5 { routedFrames.retire(make([]byte, 1000)) } if got := poolDepth(); got != framePoolBound { t.Fatalf("the pool holds %d buffers beyond its bound of %d", got, framePoolBound) } if buf := routedFrames.take(1000); buf == nil { t.Fatal("a bound-full pool answered nothing") } } func TestFramePoolAnswersNilWhenDisabled(t *testing.T) { resetPool(t) framePoolEnabled = false routedFrames.retire(make([]byte, 1000)) if got := poolDepth(); got != 0 { t.Fatalf("a disabled pool retained %d buffers", got) } if buf := routedFrames.take(1000); buf != nil { t.Fatal("a disabled pool answered a buffer") } } // TestFramePoolStaysFlatOnRepeatedRounds is the bound's proof in the // heap: the same work run again and again in one process cannot raise // the heap in use once the pool has warmed, and a rising trend is a // defect. func TestFramePoolStaysFlatOnRepeatedRounds(t *testing.T) { resetPool(t) sizes := []int64{64 << 10, 256 << 10, framePoolCeiling} round := func() { for _, s := range sizes { buf := routedFrames.take(s) if buf == nil { buf = make([]byte, s) } for i := range buf { buf[i] = byte(i) } routedFrames.retire(buf) } } for range 50 { round() } runtime.GC() var before runtime.MemStats runtime.ReadMemStats(&before) for range 200 { round() } runtime.GC() var after runtime.MemStats runtime.ReadMemStats(&after) if after.HeapInuse > before.HeapInuse+4<<20 { t.Fatalf("the heap in use rose from %d to %d bytes across 200 rounds", before.HeapInuse, after.HeapInuse) } } // routedProbeRounds is how many distinct payloads the ownership test // pushes through the hub, every one recycled through the pool. const routedProbeRounds = 300 // TestFramePoolCarriesTheBitsOverTCP is the ownership chain exercised: // a non-hub rank's frames route through the hub's pump, outbox and // drain, the drain returns every buffer the moment its write lands, // and the pump reads the next frame into what comes back, so a premature // return or a shared buffer would scramble the bits the receiver // checks, most of all under the race detector. func TestFramePoolCarriesTheBitsOverTCP(t *testing.T) { resetPool(t) runTCPWorld(t, 3, Options{Timeout: 30 * time.Second}, func(w *World) error { if w.Rank() == 2 { for i := range routedProbeRounds { buf := make([]byte, 1<<18) binary.LittleEndian.PutUint64(buf, uint64(i)) for j := 8; j < len(buf); j += 7 { buf[j] = byte(i + j) } if err := w.sendTo(1, tagHalo, buf); err != nil { return err } } return nil } if w.Rank() != 1 { return nil } for i := range routedProbeRounds { data, err := w.recvFrom(2, tagHalo) if err != nil { return err } if len(data) != 1<<18 { return base.Errf("spmd: probe %d arrived %d bytes long", i, len(data)) } if got := binary.LittleEndian.Uint64(data); got != uint64(i) { return base.Errf("spmd: probe %d arrived with the serial of %d", i, got) } for j := 8; j < len(data); j += 7 { if data[j] != byte(i+j) { return base.Errf("spmd: probe %d differs at byte %d", i, j) } } } return nil }) if poolDepth() == 0 { t.Fatal("routed traffic left the pool empty, so no buffer ever rode the chain") } }