Files

415 lines
13 KiB
Go
Raw Permalink 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 core
import (
"math"
"slices"
"strings"
"testing"
)
func maskOf(t *testing.T, a *Array) []int64 {
t.Helper()
out := make([]int64, a.Len())
for i := range out {
out[i], _ = IntAt(a, i)
}
return out
}
func TestComparisonsArray(t *testing.T) {
a := mustFromInts(t, []int64{1, 5, 5}, 3)
b := mustFromInts(t, []int64{5, 5, 9}, 3)
m, err := Lt(a, b)
if err != nil {
t.Fatalf("Lt: %v", err)
}
if got := maskOf(t, m); got[0] != 1 || got[1] != 0 || got[2] != 1 {
t.Fatalf("Lt: %v", got)
}
m, _ = Eq(a, b)
if got := maskOf(t, m); got[0] != 0 || got[1] != 1 || got[2] != 0 {
t.Fatalf("Eq: %v", got)
}
m, _ = Ne(a, b)
if got := maskOf(t, m); got[0] != 1 || got[1] != 0 || got[2] != 1 {
t.Fatalf("Ne: %v", got)
}
m, _ = Ge(a, b)
if got := maskOf(t, m); got[0] != 0 || got[1] != 1 || got[2] != 0 {
t.Fatalf("Ge: %v", got)
}
// Mixed dtypes compare exactly: int against float widens per element
// without rounding the integer side through float64 first.
f := mustFromFloats(t, []float64{0.5, 5.0, 5.5}, 3)
m, _ = Gt(a, f)
if got := maskOf(t, m); got[0] != 1 || got[1] != 0 || got[2] != 0 {
t.Fatalf("Gt mixed: %v", got)
}
if _, err := Lt(a, mustFromInts(t, []int64{1}, 1)); err == nil || !strings.Contains(err.Error(), "shape mismatch") {
t.Fatalf("comparison shape: %v", err)
}
}
func TestComparisonsScalar(t *testing.T) {
a := mustFromInts(t, []int64{1, 5, 9}, 3)
gt, err := GtI(a, 4)
if err != nil {
t.Fatalf("GtI: %v", err)
}
if got := maskOf(t, gt); got[0] != 0 || got[1] != 1 || got[2] != 1 {
t.Fatalf("GtI: %v", got)
}
le, _ := LeI(a, 5)
if got := maskOf(t, le); got[0] != 1 || got[1] != 1 || got[2] != 0 {
t.Fatalf("LeI: %v", got)
}
eq, _ := EqI(a, 5)
if got := maskOf(t, eq); got[1] != 1 || got[0] != 0 {
t.Fatalf("EqI: %v", got)
}
f := mustFromFloats(t, []float64{0.5, 2.5}, 2)
lt, _ := LtF(f, 2.0)
if got := maskOf(t, lt); got[0] != 1 || got[1] != 0 {
t.Fatalf("LtF: %v", got)
}
gtf, _ := GtI(f, 0)
if got := maskOf(t, gtf); got[0] != 1 || got[1] != 1 {
t.Fatalf("GtI on float: %v", got)
}
// IEEE NaN semantics: everything false except Ne.
nan := mustFromFloats(t, []float64{math.NaN()}, 1)
nanEq, _ := Eq(nan, nan)
if got := maskOf(t, nanEq); got[0] != 0 {
t.Fatalf("NaN Eq: %v", got)
}
nanLt, _ := Lt(nan, nan)
if got := maskOf(t, nanLt); got[0] != 0 {
t.Fatalf("NaN Lt: %v", got)
}
nanNe, _ := Ne(nan, nan)
if got := maskOf(t, nanNe); got[0] != 1 {
t.Fatalf("NaN Ne: %v", got)
}
nanEqF, _ := EqF(nan, math.NaN())
if got := maskOf(t, nanEqF); got[0] != 0 {
t.Fatalf("NaN EqF: %v", got)
}
}
func TestMaskSelect(t *testing.T) {
a := mustFromInts(t, []int64{1, 2, 3, 4}, 4)
mask, err := GtI(a, 2)
if err != nil {
t.Fatalf("GtI: %v", err)
}
selected, err := Select(a, mask)
if err != nil {
t.Fatalf("Mask: %v", err)
}
want := mustFromInts(t, []int64{3, 4}, 2)
if !Equal(want, selected) {
t.Fatalf("Mask: %s", selected)
}
// A float array keeps its dtype through masking.
f := mustFromFloats(t, []float64{0.5, 1.5, 2.5}, 3)
fm, err := GeF(f, 1.0)
if err != nil {
t.Fatalf("GeF: %v", err)
}
fs, err := Select(f, fm)
if err != nil || fs.Dtype() != Float || fs.Len() != 2 {
t.Fatalf("Mask float: %s %v", fs, err)
}
// Masks compose through the bool logic operations: And is logical
// and, Or logical or.
m1, err := GtI(a, 2) // [false, false, true, true]
if err != nil {
t.Fatalf("GtI: %v", err)
}
m2, err := GeI(a, 2) // [false, true, true, true]
if err != nil {
t.Fatalf("GeI: %v", err)
}
if m1.Dtype() != Bool || m2.Dtype() != Bool {
t.Fatalf("comparisons answered %s and %s, want bool", m1.Dtype(), m2.Dtype())
}
both, err := And(m1, m2)
if err != nil {
t.Fatalf("Mask compose: %v", err)
}
andSelected, _ := Select(a, both)
if !Equal(mustFromInts(t, []int64{3, 4}, 2), andSelected) {
t.Fatalf("Mask and: %s", andSelected)
}
either, err := Or(m1, m2)
if err != nil {
t.Fatalf("Mask compose or: %v", err)
}
orSelected, _ := Select(a, either)
if !Equal(mustFromInts(t, []int64{2, 3, 4}, 3), orSelected) {
t.Fatalf("Mask or: %s", orSelected)
}
if _, err := Select(a, f); err == nil || !strings.Contains(err.Error(), "must be a bool or int array") {
t.Fatalf("Mask dtype: %v", err)
}
if _, err := Select(a, mustFromInts(t, []int64{1}, 1)); err == nil || !strings.Contains(err.Error(), "shape mismatch") {
t.Fatalf("Mask shape: %v", err)
}
}
func TestWhere(t *testing.T) {
cond := mustFromInts(t, []int64{1, 0, 1}, 3)
x := mustFromInts(t, []int64{10, 20, 30}, 3)
y := mustFromInts(t, []int64{-1, -2, -3}, 3)
out, err := Where(cond, x, y)
if err != nil {
t.Fatalf("Where: %v", err)
}
want := mustFromInts(t, []int64{10, -2, 30}, 3)
if !Equal(want, out) {
t.Fatalf("Where: %s", out)
}
// Mixed dtypes promote to float.
fy := mustFromFloats(t, []float64{-0.5, -0.5, -0.5}, 3)
out, err = Where(cond, x, fy)
if err != nil || out.Dtype() != Float {
t.Fatalf("Where promote: %s %v", out, err)
}
if v, _ := FloatAt(out, 1); v != -0.5 {
t.Fatalf("Where promote value: %v", v)
}
if _, err := Where(fy, x, y); err == nil || !strings.Contains(err.Error(), "condition must be a bool or int array") {
t.Fatalf("Where cond dtype: %v", err)
}
if _, err := Where(mustFromInts(t, []int64{1, 0}, 2), x, y); err == nil || !strings.Contains(err.Error(), "must agree") {
t.Fatalf("Where shape: %v", err)
}
}
// TestFloatScalarComparisons pins all six relations of the scalar-float
// kernels against a hand-computed mask. The sample holds an element
// equal to the scalar, so both the <= and the >= boundary are read, and
// the float32 payload takes the same kernel one width down.
func TestFloatScalarComparisons(t *testing.T) {
f := mustFromFloats(t, []float64{-1, 0, 1, 2, 2.5, 3}, 2, 3)
cases := []struct {
name string
got func() (*Array, error)
want []bool
}{
{"LtF", func() (*Array, error) { return LtF(f, 2) }, []bool{true, true, true, false, false, false}},
{"LeF", func() (*Array, error) { return LeF(f, 2) }, []bool{true, true, true, true, false, false}},
{"GtF", func() (*Array, error) { return GtF(f, 2) }, []bool{false, false, false, false, true, true}},
{"GeF", func() (*Array, error) { return GeF(f, 2) }, []bool{false, false, false, true, true, true}},
{"EqF", func() (*Array, error) { return EqF(f, 2) }, []bool{false, false, false, true, false, false}},
{"NeF", func() (*Array, error) { return NeF(f, 2) }, []bool{true, true, true, false, true, true}},
}
for _, tc := range cases {
m, err := tc.got()
if err != nil {
t.Fatalf("%s: %v", tc.name, err)
}
if m.Dtype() != Bool {
t.Fatalf("%s answered dtype %s, want bool", tc.name, m.Dtype())
}
got := m.RawBools()
if len(got) != len(tc.want) {
t.Fatalf("%s: %d elements, want %d", tc.name, len(got), len(tc.want))
}
for i := range tc.want {
if got[i] != tc.want[i] {
t.Fatalf("%s over [-1 0 1 2 2.5 3] against 2: mask %v, want %v",
tc.name, got, tc.want)
}
}
}
// The float32 payload one width down, boundary included.
g := mustFromFloat32s(t, []float32{-1, 2, 2.5}, 3)
le, err := LeF(g, 2)
if err != nil {
t.Fatalf("LeF float32: %v", err)
}
if got := le.RawBools(); got[0] != true || got[1] != true || got[2] != false {
t.Fatalf("LeF float32 over [-1 2 2.5] against 2: %v, want [true true false]", got)
}
gt, err := GtF(g, 2)
if err != nil {
t.Fatalf("GtF float32: %v", err)
}
if got := gt.RawBools(); got[0] != false || got[1] != false || got[2] != true {
t.Fatalf("GtF float32 over [-1 2 2.5] against 2: %v, want [false false true]", got)
}
}
// TestWhereDenseFloatOperands pins the operand order of the dense pick
// over two float payloads of one width: a nonzero condition takes the
// first operand, a zero the second, on both the float64 and the float32
// kernel.
func TestWhereDenseFloatOperands(t *testing.T) {
cond := mustFromInts(t, []int64{1, 0, 1, 0, 1}, 5)
x := mustFromFloats(t, []float64{10, 20, 30, 40, 50}, 5)
y := mustFromFloats(t, []float64{-1, -2, -3, -4, -5}, 5)
out, err := Where(cond, x, y)
if err != nil {
t.Fatalf("Where float: %v", err)
}
if out.Dtype() != Float {
t.Fatalf("Where float dtype: %s", out.Dtype())
}
for i, want := range []float64{10, -2, 30, -4, 50} {
if got := out.RawFloats()[i]; got != want {
t.Fatalf("Where float [%d] = %v, want %v", i, got, want)
}
}
// An all-zero and an all-one condition pin each arm in full.
zeros := mustFromInts(t, []int64{0, 0, 0, 0, 0}, 5)
low, err := Where(zeros, x, y)
if err != nil {
t.Fatalf("Where float zeros: %v", err)
}
for i := range 5 {
if got := low.RawFloats()[i]; got != y.RawFloats()[i] {
t.Fatalf("Where float zero condition [%d] = %v, want the second operand %v",
i, got, y.RawFloats()[i])
}
}
ones := mustFromInts(t, []int64{1, 1, 1, 1, 1}, 5)
high, err := Where(ones, x, y)
if err != nil {
t.Fatalf("Where float ones: %v", err)
}
for i := range 5 {
if got := high.RawFloats()[i]; got != x.RawFloats()[i] {
t.Fatalf("Where float one condition [%d] = %v, want the first operand %v",
i, got, x.RawFloats()[i])
}
}
// The float32 pair takes its own kernel.
c32 := mustFromInts(t, []int64{0, 1, 1}, 3)
x32 := mustFromFloat32s(t, []float32{1.5, -2.5, 7.25}, 3)
y32 := mustFromFloat32s(t, []float32{9, 9, 9}, 3)
out32, err := Where(c32, x32, y32)
if err != nil {
t.Fatalf("Where float32: %v", err)
}
for i, want := range []float32{9, -2.5, 7.25} {
if got := out32.RawFloat32s()[i]; got != want {
t.Fatalf("Where float32 [%d] = %v, want %v", i, got, want)
}
}
}
// TestLogicOpsBoolContract pins And, Or, Xor and Not: bool operands
// only, element-wise boolean semantics, bool results, and a loud named
// refusal for every other dtype.
func TestLogicOpsBoolContract(t *testing.T) {
x := narrowBools(t, []bool{true, true, false, false}, 4)
y := narrowBools(t, []bool{true, false, true, false}, 4)
check := func(name string, got *Array, err error, want []bool) {
t.Helper()
if err != nil {
t.Fatalf("%s: %v", name, err)
}
if got.Dtype() != Bool {
t.Fatalf("%s answered dtype %s, want bool", name, got.Dtype())
}
if !slices.Equal(got.RawBools(), want) {
t.Fatalf("%s = %v, want %v", name, got.RawBools(), want)
}
}
and, err := And(x, y)
check("And", and, err, []bool{true, false, false, false})
or, err := Or(x, y)
check("Or", or, err, []bool{true, true, true, false})
xor, err := Xor(x, y)
check("Xor", xor, err, []bool{false, true, true, false})
not, err := Not(x)
check("Not", not, err, []bool{false, false, true, true})
// The refusals name the actual dtypes.
ints := mustFromInts(t, []int64{1, 0, 1, 0}, 4)
if _, err := And(x, ints); err == nil ||
!strings.Contains(err.Error(), "operands must be bool arrays, got bool and int") {
t.Fatalf("And with an int operand: %v", err)
}
if _, err := Not(ints); err == nil ||
!strings.Contains(err.Error(), "operands must be bool arrays, got int") {
t.Fatalf("Not on an int operand: %v", err)
}
short := narrowBools(t, []bool{true}, 1)
if _, err := Or(x, short); err == nil || !strings.Contains(err.Error(), "shape mismatch") {
t.Fatalf("Or with a shape mismatch: %v", err)
}
}
// TestWhereBoolCondition pins Where's condition contract: an int mask or
// a bool condition, the pinned refusal wording for every other dtype,
// and promoted result dtypes written correctly, narrow ones included.
func TestWhereBoolCondition(t *testing.T) {
cond := narrowBools(t, []bool{true, false, true, false}, 4)
x8 := narrowInt8s(t, []int8{1, 2, 3, 4}, 4)
y8 := narrowInt8s(t, []int8{9, 9, 9, 9}, 4)
out, err := Where(cond, x8, y8)
if err != nil {
t.Fatalf("Where with a bool condition: %v", err)
}
if out.Dtype() != Int8 {
t.Fatalf("Where bool-cond int8 answered %s, want int8", out.Dtype())
}
if want := []int8{1, 9, 3, 9}; !slices.Equal(out.RawInt8s(), want) {
t.Fatalf("Where bool-cond int8 = %v, want %v", out.RawInt8s(), want)
}
// A mixed integer-class pair promotes through the containment table
// and the fallback walk writes the promoted payload.
u8y := narrowUint8s(t, []uint8{9, 9, 9, 9}, 4)
icond := mustFromInts(t, []int64{1, 0, 1, 0}, 4)
mix, err := Where(icond, x8, u8y)
if err != nil {
t.Fatalf("Where int8 with uint8: %v", err)
}
if mix.Dtype() != Int16 {
t.Fatalf("Where int8 with uint8 answered %s, want int16", mix.Dtype())
}
if want := []int16{1, 9, 3, 9}; !slices.Equal(mix.RawInt16s(), want) {
t.Fatalf("Where int8 with uint8 = %v, want %v", mix.RawInt16s(), want)
}
// A bool condition over bool operands answers bool.
xb := narrowBools(t, []bool{true, false, true, false}, 4)
yb := narrowBools(t, []bool{false, false, false, false}, 4)
bb, err := Where(cond, xb, yb)
if err != nil {
t.Fatalf("Where bool over bool: %v", err)
}
if bb.Dtype() != Bool || !slices.Equal(bb.RawBools(), []bool{true, false, true, false}) {
t.Fatalf("Where bool over bool = %s %v", bb.Dtype(), bb.RawBools())
}
// Every other condition dtype keeps the pinned refusal wording.
fc := mustFromFloats(t, []float64{1, 0, 1, 0}, 4)
if _, err := Where(fc, x8, y8); err == nil ||
!strings.Contains(err.Error(), "condition must be a bool or int array") {
t.Fatalf("Where with a float condition: %v", err)
}
}