// Copyright (c) 2026 Petr BalvĂ­n (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) } }