// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) // SPDX-License-Identifier: MIT package core import ( "runtime" "testing" ) func TestSetNumCPU(t *testing.T) { // SetNumCPU(0) resets to runtime.NumCPU(): the invariant that must // hold regardless of what previous tests set. Test against that, // not against a captured "previous" value (tests run in any order). ncpu := runtime.NumCPU() if NumWorkers() < 1 { t.Fatalf("NumWorkers: %d", NumWorkers()) } got := SetNumCPU(4) if got < 1 { t.Errorf("SetNumCPU return: %d, want ≥ 1", got) } if NumWorkers() != 4 { t.Errorf("NumWorkers after SetNumCPU(4): %d", NumWorkers()) } // workersFor bounds against the explicit worker count, independent // of the host's CPU count (the CI runner has one core). if w := workersFor(2); w != 2 { t.Errorf("workersFor(2) with 4 workers: %d, want 2", w) } if w := workersFor(1); w != 1 { t.Errorf("workersFor(1): %d, want 1", w) } if w := workersFor(10); w != 4 { t.Errorf("workersFor(10) with 4 workers: %d, want 4", w) } SetNumCPU(0) // resets to NumCPU if NumWorkers() != ncpu { t.Errorf("NumWorkers after reset: %d, want %d", NumWorkers(), ncpu) } // Universal invariants after the reset: at least one worker, never // more than the item count. if w := workersFor(0); w != 1 { t.Errorf("workersFor(0): %d, want 1 (floor of one worker)", w) } if w := workersFor(1); w != 1 { t.Errorf("workersFor(1) after reset: %d, want 1", w) } } // TestParallelCoverage forces the parallel branch of `parallel` to run // even on a single-core CI runner, where workersFor(n) would otherwise // collapse to 1 and skip the goroutine-spawning path entirely, which // silently drops the measured coverage below the gate. It also covers // the parallel-only merge branch of the axis reductions (reduceAxis), // whose private-scratch/merge code has no serial equivalent. func TestParallelCoverage(t *testing.T) { prev := SetNumCPU(2) defer SetNumCPU(prev) a, err := FromFloats(make([]float64, 1<<16), 1<<16) if err != nil { t.Fatal(err) } b, err := FromFloats(make([]float64, 1<<16), 1<<16) if err != nil { t.Fatal(err) } // With numWorkers=2 and 65536 items the parallel branch spawns // goroutines even on a one-core host. sum, err := Add(a, b) if err != nil { t.Fatal(err) } if sum.Len() != a.Len() { t.Fatalf("Add len: %d", sum.Len()) } // Axis reductions take the parallel merge path with 2 workers, // covering the private-scratch and merge code in reduceAxis. mat, err := Reshape(a, 256, 256) if err != nil { t.Fatal(err) } if _, err := SumAxis(mat, 1); err != nil { t.Fatal(err) } if _, err := MaxAxis(mat, 0); err != nil { t.Fatal(err) } }