// Copyright (c) 2026 Petr BalvĂ­n (https://petrbalvin.org) // SPDX-License-Identifier: MIT // Regression pins for input validation in internal/core: one test per // defect, named after what it pins; the radix chunk split also carries // a deterministic invariant test. package core import ( "fmt" "math" "slices" "strings" "sync" "testing" "time" "sourcedock.dev/petrbalvin/tensor/internal/engine" ) // callNoPanic runs fn on the test goroutine and reports a panic as a // test failure: a validation gap must surface as an error the caller can // handle, never as a crash inside the library. func callNoPanic(t *testing.T, what string, fn func() error) error { t.Helper() var err error func() { defer func() { if r := recover(); r != nil { t.Fatalf("%s panicked instead of returning an error: %v", what, r) } }() err = fn() }() return err } // Pad used to accept negative pad values, drive the new shape // negative and die in alloc with "makeslice: len out of range". Pad // documents an error for a malformed pad argument, so the negatives have // to be refused by name. func TestPadRejectsNegativePadValues(t *testing.T) { a := mustFromFloats(t, []float64{1, 2, 3}, 3) err := callNoPanic(t, "Pad(-2, -2)", func() error { _, err := Pad(a, []int{-2, -2}, "constant", 0) return err }) if err == nil { t.Fatal("Pad accepted the negative pad pair (-2, -2)") } if !strings.Contains(err.Error(), "non-negative") { t.Fatalf("Pad error does not name the negative pads: %v", err) } // One negative side of a pair is just as malformed. err = callNoPanic(t, "Pad(1, -1)", func() error { _, err := Pad(a, []int{1, -1}, "constant", 0) return err }) if err == nil { t.Fatal("Pad accepted the mixed pad pair (1, -1)") } // The valid path is untouched. ok, err := Pad(a, []int{1, 1}, "constant", 0) if err != nil { t.Fatalf("Pad valid pair: %v", err) } if !slices.Equal(ok.Shape(), []int{5}) { t.Fatalf("Pad valid pair shape %v, want [5]", ok.Shape()) } } // TruncatedNormal with std = NaN used to pass the `std <= 0` // guard, leave the rejection window NaN and spin in the retry loop // forever. The draw must fall back to the degenerate all-zero result // promptly, on this goroutine: the watchdog keeps a regression from // hanging the suite. func TestTruncatedNormalNaNStdReturnsPromptly(t *testing.T) { done := make(chan *Array, 1) go func() { done <- TruncatedNormal(NewGenerator(1), []int{3}, 0, math.NaN()) }() select { case got := <-done: if got == nil { t.Fatal("TruncatedNormal(std=NaN) returned nil") } for i, v := range got.RawFloat32s() { if v != 0 { t.Fatalf("TruncatedNormal(std=NaN) value %d = %v, want the degenerate 0", i, v) } } case <-time.After(5 * time.Second): t.Fatal("TruncatedNormal(std=NaN) still running after 5s: the rejection loop cannot exit") } // The documented degenerate path (std <= 0) still returns zeros, and // a valid std still draws inside the window. zero := TruncatedNormal(NewGenerator(2), []int{4}, 0, 0) if zero == nil { t.Fatal("TruncatedNormal(std=0) returned nil") } for i, v := range zero.RawFloat32s() { if v != 0 { t.Fatalf("TruncatedNormal(std=0) value %d = %v, want 0", i, v) } } drawn := TruncatedNormal(NewGenerator(3), []int{64}, 0, 1) if drawn == nil { t.Fatal("TruncatedNormal(std=1) returned nil") } for i, v := range drawn.RawFloat32s() { if v < -2 || v > 2 { t.Fatalf("TruncatedNormal(std=1) value %d = %v outside the +-2 sigma window", i, v) } } } // Normal let a NaN std through (`std < 0` is false for NaN) and // returned an array of NaNs silently. This is the feeder of the TruncatedNormal case above, so it // has to be a loud error. func TestNormalRejectsNaNStd(t *testing.T) { g := NewGenerator(1) if _, err := Normal(g, 3, 0, math.NaN()); err == nil { t.Fatal("Normal accepted a NaN std and drew silently wrong values") } if _, err := Normal(g, 3, 0, -1); err == nil { t.Fatal("Normal accepted a negative std") } arr, err := Normal(g, 3, 0, 1) if err != nil { t.Fatalf("Normal std=1: %v", err) } for i, v := range arr.RawFloats() { if math.IsNaN(v) { t.Fatalf("Normal std=1 value %d is NaN", i) } } } // InterpolateGrid with a NaN query used to pass both clamps, // convert to the platform's indefinite integer and panic inside // FloatAt; with a stride above 1 it silently returned a value read from // a wrapped index instead. The documented clamp must hold or the call // must be refused. +Inf and -Inf keep clamping, as documented. func TestInterpolateGridRejectsNaNQuery(t *testing.T) { grid := mustFloats(t, []float64{0, 10, 20, 30}, 4) origins := []float64{0} steps := []float64{1} nan := mustFloats(t, []float64{math.NaN()}, 1, 1) err := callNoPanic(t, "InterpolateGrid(NaN)", func() error { _, err := InterpolateGrid(grid, origins, steps, nan) return err }) if err == nil { t.Fatal("InterpolateGrid accepted a NaN query") } if !strings.Contains(err.Error(), "NaN") { t.Fatalf("InterpolateGrid error does not name the NaN position: %v", err) } // The infinite queries keep their clamp: -Inf to the first sample, // +Inf to the last. inf := mustFloats(t, []float64{math.Inf(-1), math.Inf(1)}, 2, 1) out, err := InterpolateGrid(grid, origins, steps, inf) if err != nil { t.Fatalf("InterpolateGrid with infinite queries: %v", err) } if got := out.FloatAt(0); got != 0 { t.Fatalf("-Inf query = %v, want the first sample 0", got) } if got := out.FloatAt(1); got != 30 { t.Fatalf("+Inf query = %v, want the last sample 30", got) } } // HaltonPoints had no skip + n bound, so an index near MaxInt // wrapped the int arithmetic and every point collapsed to the origin // with no error. The constructor now enforces the same 2^32 index // budget SobolPoints does. func TestHaltonPointsRejectsSkipOverflow(t *testing.T) { if _, err := HaltonPoints(2, 2, math.MaxInt); err == nil { t.Fatal("HaltonPoints accepted skip = MaxInt: the points collapse to the origin silently") } // The boundary the twin constructor enforces: n + skip must stay // below 2^32, so the last acceptable index is 2^32 - 1. if _, err := HaltonPoints(2, 2, 1<<32-2); err == nil { t.Fatal("HaltonPoints accepted n + skip = 2^32") } if _, err := HaltonPoints(2, 2, 1<<32-3); err != nil { t.Fatalf("HaltonPoints rejected n + skip = 2^32 - 1: %v", err) } if _, err := HaltonPoints(1<<32, 2, 0); err == nil { t.Fatal("HaltonPoints accepted n = 2^32") } // The ordinary path is untouched. if _, err := HaltonPoints(4, 2, 3); err != nil { t.Fatalf("HaltonPoints valid skip: %v", err) } } // Unsqueeze(-1) appends the new axis at the end (the convention // the code has always followed and the frozen API keeps); the doc // comment claimed it inserts before the last dimension. The behaviour // is pinned here so the comment and the code cannot drift apart again // in either direction. func TestUnsqueezeNegativeDimAppendsAxis(t *testing.T) { a := mustFromInts(t, []int64{1, 2, 3, 4, 5, 6}, 2, 3) got, err := Unsqueeze(a, -1) if err != nil { t.Fatalf("Unsqueeze(-1): %v", err) } if !slices.Equal(got.Shape(), []int{2, 3, 1}) { t.Fatalf("Unsqueeze((2,3), -1) shape %v, want the appended [2 3 1]", got.Shape()) } // Counting from the end of the result rank: -2 lands one axis earlier. got, err = Unsqueeze(a, -2) if err != nil { t.Fatalf("Unsqueeze(-2): %v", err) } if !slices.Equal(got.Shape(), []int{2, 1, 3}) { t.Fatalf("Unsqueeze((2,3), -2) shape %v, want [2 1 3]", got.Shape()) } // One step past the rank is refused. if _, err := Unsqueeze(a, -4); err == nil { t.Fatal("Unsqueeze(-4) accepted for a rank-2 input") } } // seekCoord was dead code with no caller in the module. Removing // it must not move a single element, so BroadcastTo is pinned against // the naive per-element odometer reference the run-filled fast path // replaced. func TestBroadcastToMatchesOdometerReference(t *testing.T) { cases := []struct { name string vals []int64 srcSh []int target []int }{ {"column", []int64{1, 2, 3}, []int{3, 1}, []int{3, 2}}, {"prepend", []int64{1, 2, 3}, []int{3, 1}, []int{2, 3, 1}}, {"scalar", []int64{7}, []int{1}, []int{2, 3, 4}}, {"bias", []int64{1, 2, 3, 4, 5, 6}, []int{1, 3, 2}, []int{4, 3, 2}}, {"same-shape", []int64{9, 8, 7, 6}, []int{2, 2}, []int{2, 2}}, {"trailing-run", []int64{1, 2}, []int{2, 1}, []int{2, 3}}, } for _, tc := range cases { t.Run(tc.name, func(t *testing.T) { src := mustFromInts(t, tc.vals, tc.srcSh...) got, err := BroadcastTo(src, tc.target...) if err != nil { t.Fatalf("BroadcastTo: %v", err) } want := broadcastOdometerReference(t, src, tc.target) if !slices.Equal(want, got.RawInts()) { t.Fatalf("BroadcastTo diverged from the odometer reference:\n got %v\nwant %v", got.RawInts(), want) } }) } } // broadcastOdometerReference expands src to target one element at a // time, coordinates recomputed from scratch: the slow path BroadcastTo's // run fill replaced. func broadcastOdometerReference(t *testing.T, src *Array, target []int) []int64 { t.Helper() total := 1 for _, d := range target { total *= d } sh := src.Shape() out := make([]int64, total) coord := make([]int, len(target)) srcCoord := make([]int, len(sh)) off := len(target) - len(sh) for i := range total { for d := range sh { c := 0 if sh[d] != 1 { c = coord[off+d] } srcCoord[d] = c } v, err := IntAt(src, srcCoord...) if err != nil { t.Fatalf("IntAt(%v): %v", srcCoord, err) } out[i] = v advanceOdometer(coord, target) } return out } // Radix chunk split, invariant half: the parallel radix derives its histogram rows // from its own split of [0, n) instead of re-reading the global worker // count the way engine.ParallelMin did, so a concurrent SetNumCPU can // move work between goroutines but never two live chunks onto one row. // This test reproduces the disagreement deterministically and with no // reliance on scheduling: the stale-snapshot case raises the live worker // count while the chunk width handed to the split stays the one a caller // would have snapshotted earlier, which is exactly the window SetNumCPU // opens (the raising read is what used to collapse two live chunks onto // one row). Every row must be handed out once with the range it owns, // start = row*chunk, and the ranges must partition [0, n); the remaining // cases pin that contract either side of the radixParMin spawn floor. func TestRadixChunksRowMatchesOwnedRange(t *testing.T) { cases := []struct { name string n int workers int chunk int // the caller's snapshot width, stale where noted }{ {"exact-fit", 40_000, 2, 20_000}, {"ragged", 40_001, 3, 13_334}, {"one-chunk", 40_000, 1, 40_000}, {"floor-exact", 40_000, 1, radixParMin}, {"stale-snapshot", 40_000, 8, 20_000}, // width from a 2-worker snapshot {"below-floor", 40_000, 8, radixParMin - 1}, } for _, tc := range cases { t.Run(tc.name, func(t *testing.T) { prev := engine.SetNumWorkers(tc.workers) defer engine.SetNumWorkers(prev) parallel := tc.chunk >= radixParMin && tc.chunk < tc.n nchunks := 1 if parallel { nchunks = (tc.n + tc.chunk - 1) / tc.chunk } // The callback may run on a spawned goroutine, so failures // are collected and reported on the test goroutine. var mu sync.Mutex var problems []string touched := make([]int, tc.n) rows := make([]bool, nchunks) forEachRadixChunk(tc.n, tc.chunk, func(row, start, end int) { var local []string if row < 0 || row >= nchunks { local = append(local, fmt.Sprintf("row %d out of range for %d chunks", row, nchunks)) } else if rows[row] { local = append(local, fmt.Sprintf("row %d handed out twice", row)) } else { rows[row] = true } // The row identity is the chunk position: the range a // worker owns and the histogram row it writes are the // same object, so they cannot disagree. wantStart, wantEnd := 0, tc.n if parallel { wantStart, wantEnd = row*tc.chunk, min(row*tc.chunk+tc.chunk, tc.n) } if start != wantStart || end != wantEnd { local = append(local, fmt.Sprintf("row %d owns [%d, %d), want [%d, %d)", row, start, end, wantStart, wantEnd)) if start < 0 || end > tc.n || start >= end { local = append(local, fmt.Sprintf("row %d owns the illegal range [%d, %d)", row, start, end)) } } if start >= 0 && end <= tc.n && start < end { for i := start; i < end; i++ { touched[i]++ } } mu.Lock() problems = append(problems, local...) mu.Unlock() }) if len(problems) > 0 { t.Fatalf("chunk split broken: %s", strings.Join(problems, "; ")) } for i, c := range touched { if c != 1 { t.Fatalf("index %d visited %d times", i, c) } } for row, seen := range rows { if !seen { t.Fatalf("row %d never ran", row) } } }) } } // flipWorkers runs fn at least rounds times while a second goroutine // flips the global worker count between low and high. SetNumCPU is // documented as safe to call at any time, so a kernel that derives its // chunking from the global count must still finish correctly; before // Radix chunk split, growth inside the kernel's own window collapsed two live chunks // onto one histogram row and corrupted the output. func flipWorkers(t *testing.T, low, high, rounds int, fn func(round int)) { t.Helper() prev := engine.SetNumWorkers(low) defer engine.SetNumWorkers(prev) stop := make(chan struct{}) var flipper sync.WaitGroup flipper.Go(func() { w := low for { select { case <-stop: return default: } if w == low { w = high } else { w = low } engine.SetNumWorkers(w) } }) defer func() { close(stop) flipper.Wait() }() for round := range rounds { fn(round) } } // Radix chunk split, value-radix half: Sort of a large float64 payload must stay an // ascending permutation while a concurrent SetNumCPU moves the worker // count from 2 to 8 and back. The output is checked against the sorted // order itself, so a lost histogram update (rows colliding), a // mis-ordered scatter (bit-identity broken across chunks) or an // out-of-range write all fail here. func TestSortRadixStableUnderConcurrentSetNumCPU(t *testing.T) { const n = 200_000 g := NewGenerator(7) src, err := Floats(g, n) if err != nil { t.Fatalf("Floats: %v", err) } template := slices.Clone(src.RawFloats()) // Ties force the scatter to be exercised across chunk boundaries. for i := range template { if i%3 == 0 { template[i] = float64(i % 17) } } reference := sortReference(template) flipWorkers(t, 2, 8, 8, func(round int) { a := mustFromFloats(t, slices.Clone(template), n) got, err := Sort(a) if err != nil { t.Fatalf("round %d: Sort: %v", round, err) } vals := got.RawFloats() if !equalSortedFloats(reference, vals) { t.Fatalf("round %d: Sort corrupted the parallel radix output", round) } }) } // Radix chunk split, permutation-radix half: ArgSort must keep returning the same // stable permutation the serial radix returns while the worker count // changes concurrently. A row collision shows up either as a repeated // or missing index or as a change of the permutation itself. func TestArgSortRadixStableUnderConcurrentSetNumCPU(t *testing.T) { const n = 200_000 g := NewGenerator(11) src, err := Floats(g, n) if err != nil { t.Fatalf("Floats: %v", err) } vals := slices.Clone(src.RawFloats()) for i := range vals { if i%3 == 0 { vals[i] = float64(i % 17) } } reference := argSortReference(vals) flipWorkers(t, 2, 8, 8, func(round int) { a := mustFromFloats(t, slices.Clone(vals), n) got, err := ArgSort(a) if err != nil { t.Fatalf("round %d: ArgSort: %v", round, err) } if !slices.Equal(reference, got.RawInts()) { t.Fatalf("round %d: ArgSort diverged from the stable serial permutation", round) } }) }