Files
tensor/spmd/framepool_test.go
T

213 lines
6.0 KiB
Go
Raw 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"
"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")
}
}