// Copyright (c) 2026 Petr BalvĂ­n (https://petrbalvin.org) // SPDX-License-Identifier: MIT package core import ( "math" "slices" "strings" "testing" ) // Pins for the parallel payload walks: strided operands must be read in // logical order through the accessors, Take must report the smallest // offending position, and the radix ArgSort permutation must equal the // stable comparator permutation on tie-heavy inputs. func pinFloats(t *testing.T, vals []float64, shape ...int) *Array { t.Helper() a, err := FromFloats(vals, shape...) if err != nil { t.Fatal(err) } return a } func pinInts(t *testing.T, vals []int64, shape ...int) *Array { t.Helper() a := New(Int, shape...) copy(a.RawInts(), vals) return a } // pinStridedFloat builds a test-mechanics strided array: logical (r, c) // reads payload[r*rowStride + c]. The payload carries an invisible tail // element at an unaddressed slot, the state a raw payload walk gets // wrong. func pinStridedFloat(t *testing.T, payload []float64, shape []int, strides []int) *Array { t.Helper() return &Array{shape: shape, dt: Float, floats: payload, strides: strides} } func pinStridedInt(t *testing.T, payload []int64, shape []int, strides []int) *Array { t.Helper() return &Array{shape: shape, dt: Int, ints: payload, strides: strides} } func TestArgwhereNonzeroStridedUseLogicalElements(t *testing.T) { // Logical window [[1, 0], [0, 5]] at strides [3, 1]: the elements // sit at payload slots 0, 1, 3, 4, and slot 2 holds an invisible // 99. A payload walk at the logical position reads slot 3 for the // bottom-right element and calls the 5 a zero. stridedF := pinStridedFloat(t, []float64{1, 0, 99, 0, 5, 7, 7, 7}, []int{2, 2}, []int{3, 1}) twinF := pinFloats(t, []float64{1, 0, 0, 5}, 2, 2) gotF, err := Argwhere(stridedF) if err != nil { t.Fatalf("Argwhere(strided float): %v", err) } wantF, err := Argwhere(twinF) if err != nil { t.Fatalf("Argwhere(twin float): %v", err) } if !Equal(gotF, wantF) { t.Fatalf("Argwhere(strided float) = %v, want %v", gotF, wantF) } nzS, err := Nonzero(stridedF) if err != nil { t.Fatalf("Nonzero(strided float): %v", err) } nzT, err := Nonzero(twinF) if err != nil { t.Fatalf("Nonzero(twin float): %v", err) } if !slices.Equal(nzS[0], nzT[0]) || !slices.Equal(nzS[1], nzT[1]) { t.Fatalf("Nonzero(strided float) = %v, want %v", nzS, nzT) } stridedI := pinStridedInt(t, []int64{1, 0, 99, 0, 5, 7, 7, 7}, []int{2, 2}, []int{3, 1}) twinI := pinInts(t, []int64{1, 0, 0, 5}, 2, 2) gotI, err := Argwhere(stridedI) if err != nil { t.Fatalf("Argwhere(strided int): %v", err) } wantI, err := Argwhere(twinI) if err != nil { t.Fatalf("Argwhere(twin int): %v", err) } if !Equal(gotI, wantI) { t.Fatalf("Argwhere(strided int) = %v, want %v", gotI, wantI) } nzSI, err := Nonzero(stridedI) if err != nil { t.Fatalf("Nonzero(strided int): %v", err) } nzTI, err := Nonzero(twinI) if err != nil { t.Fatalf("Nonzero(twin int): %v", err) } if !slices.Equal(nzSI[0], nzTI[0]) || !slices.Equal(nzSI[1], nzTI[1]) { t.Fatalf("Nonzero(strided int) = %v, want %v", nzSI, nzTI) } } func TestTakeReportsSmallestOffender(t *testing.T) { src := pinFloats(t, []float64{10, 11, 12, 13, 14}, 5) idx := pinInts(t, []int64{99, 88, 0}, 3) _, err := Take(src, idx) if err == nil { t.Fatal("Take accepted out-of-range indices") } if !strings.Contains(err.Error(), "position 0") || strings.Contains(err.Error(), "position 1") { t.Fatalf("Take offender report = %q, want the smallest offending position 0", err.Error()) } } func TestArgSortTiePermutationsMatchStableReference(t *testing.T) { for _, n := range []int{64, 1024, 4096} { vals := make([]float64, n) ivals := make([]int64, n) state := uint64(0x9E3779B97F4A7C15 + uint64(n)) for i := range vals { state = state*6364136223846793005 + 1442695040888963407 // A small level set, so ties dominate the permutation. level := int64(state>>60) % 5 ivals[i] = level - 2 vals[i] = float64(level - 2) } ref := make([]int, n) for i := range ref { ref[i] = i } // The stable comparator permutation: ties keep index order, the // contract a stable digit scatter must reproduce. slices.SortStableFunc(ref, func(x, y int) int { switch { case vals[x] < vals[y]: return -1 case vals[x] > vals[y]: return 1 } return 0 }) gotF, err := ArgSort(pinFloats(t, vals, n)) if err != nil { t.Fatalf("n=%d ArgSort(float): %v", n, err) } gotI, err := ArgSort(pinInts(t, ivals, n)) if err != nil { t.Fatalf("n=%d ArgSort(int): %v", n, err) } // ArgSort answers the permutation as an int array whatever the // sorted dtype was. rf, ri := gotF.RawInts(), gotI.RawInts() for i := range ref { if int(rf[i]) != ref[i] { t.Fatalf("n=%d float permutation at %d = %d, want %d (values %v)", n, i, int(rf[i]), ref[i], vals) } if int(ri[i]) != ref[i] { t.Fatalf("n=%d int permutation at %d = %d, want %d (values %v)", n, i, int(ri[i]), ref[i], ivals) } } } // The NaN and signed-zero contract on the same walk: NaN sorts // last, -0 folds onto +0 in the value order but keeps its index // order among the zeros. withSpecial := pinFloats(t, []float64{2, math.NaN(), math.Copysign(0, -1), -1, 0, math.NaN(), 3}, 7) got, err := ArgSort(withSpecial) if err != nil { t.Fatalf("ArgSort(special): %v", err) } perm := got.RawInts() // Sorted values: -1, then the zeros at indices 2 and 4 in index // order, then 2, 3, then the NaNs at indices 1 and 5 in index order. want := []int{3, 2, 4, 0, 6, 1, 5} for i := range want { if int(perm[i]) != want[i] { t.Fatalf("special permutation at %d = %d, want %d (full %v)", i, int(perm[i]), want[i], perm) } } }