// Copyright (c) 2026 Petr BalvĂ­n (https://petrbalvin.org) // SPDX-License-Identifier: MIT package core import ( "math" "slices" "strings" "testing" ) func intsOf(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 floatsOf(t *testing.T, a *Array) []float64 { t.Helper() out := make([]float64, a.Len()) for i := range out { out[i], _ = FloatAt(a, i) } return out } func TestElementwiseIntStaysInt(t *testing.T) { a := mustFromInts(t, []int64{1, 2, 3}, 3) b := mustFromInts(t, []int64{4, 5, 6}, 3) sum, err := Add(a, b) if err != nil { t.Fatalf("Add: %v", err) } if sum.Dtype() != Int { t.Fatalf("int + int must stay int, got %s", sum.Dtype()) } want := []int64{5, 7, 9} got := intsOf(t, sum) for i := range want { if got[i] != want[i] { t.Fatalf("Add: %v", got) } } diff, _ := Sub(b, a) if diff.Dtype() != Int || intsOf(t, diff)[0] != 3 { t.Fatalf("Sub: %s %v", diff, intsOf(t, diff)) } prod, _ := Mul(a, b) if prod.Dtype() != Int || intsOf(t, prod)[2] != 18 { t.Fatalf("Mul: %s %v", prod, intsOf(t, prod)) } // int arithmetic wraps like Go's int64. big := mustFromInts(t, []int64{math.MaxInt64}, 1) one := mustFromInts(t, []int64{1}, 1) wrapped, _ := Add(big, one) if v, _ := IntAt(wrapped, 0); v != math.MinInt64 { t.Fatalf("wrap: %d", v) } } func TestElementwisePromotesToFloat(t *testing.T) { i := mustFromInts(t, []int64{1, 2}, 2) f := mustFromFloats(t, []float64{0.5, 1.5}, 2) sum, err := Add(i, f) if err != nil { t.Fatalf("Add: %v", err) } if sum.Dtype() != Float { t.Fatalf("int + float must promote, got %s", sum.Dtype()) } got := floatsOf(t, sum) if got[0] != 1.5 || got[1] != 3.5 { t.Fatalf("Add promoted: %v", got) } } func TestDivIsTrueDivision(t *testing.T) { a := mustFromInts(t, []int64{1, 7, -7}, 3) b := mustFromInts(t, []int64{2, 2, 2}, 3) q, err := Div(a, b) if err != nil { t.Fatalf("Div: %v", err) } if q.Dtype() != Float { t.Fatalf("Div must always yield float, got %s", q.Dtype()) } got := floatsOf(t, q) if got[0] != 0.5 || got[1] != 3.5 || got[2] != -3.5 { t.Fatalf("Div: %v", got) } // Zero divisors are IEEE, never errors. num := mustFromFloats(t, []float64{1.0, -1.0, 0.0}, 3) zero := mustFromFloats(t, []float64{0.0, 0.0, 0.0}, 3) zq, err := Div(num, zero) if err != nil { t.Fatalf("Div by zero: %v", err) } z := floatsOf(t, zq) if !math.IsInf(z[0], 1) || !math.IsInf(z[1], -1) || !math.IsNaN(z[2]) { t.Fatalf("IEEE: %v", z) } } func TestQuo(t *testing.T) { a := mustFromInts(t, []int64{7, -7, 9}, 3) b := mustFromInts(t, []int64{2, 2, 3}, 3) q, err := Quo(a, b) if err != nil { t.Fatalf("Quo: %v", err) } got := intsOf(t, q) if got[0] != 3 || got[1] != -3 || got[2] != 3 { t.Fatalf("Quo: %v", got) } f := mustFromFloats(t, []float64{1}, 1) if _, err := Quo(f, f); err == nil || !strings.Contains(err.Error(), "needs int arrays") { t.Fatalf("Quo float: %v", err) } z := mustFromInts(t, []int64{1, 0}, 2) o := mustFromInts(t, []int64{1, 1}, 2) if _, err := Quo(z, o); err != nil { t.Fatalf("Quo zeros in dividend is fine: %v", err) } if _, err := Quo(o, z); err == nil || !strings.Contains(err.Error(), "division by zero") { t.Fatalf("Quo zero divisor: %v", err) } } func TestShapeMismatchIsLoud(t *testing.T) { a := mustFromInts(t, []int64{1, 2, 3}, 3) b := mustFromInts(t, []int64{1, 2}, 2) _, err := Add(a, b) if err == nil || !strings.Contains(err.Error(), "shape mismatch (3) vs (2)") { t.Fatalf("shape mismatch: %v", err) } if _, err := Div(a, b); err == nil || !strings.Contains(err.Error(), "shape mismatch") { t.Fatalf("Div shape mismatch: %v", err) } if _, err := Quo(a, b); err == nil || !strings.Contains(err.Error(), "shape mismatch") { t.Fatalf("Quo shape mismatch: %v", err) } } func TestScalarOps(t *testing.T) { a := mustFromInts(t, []int64{1, 2}, 2) if v, _ := IntAt(AddI(a, 5), 0); v != 6 { t.Fatalf("AddI: %d", v) } if v, _ := IntAt(SubI(a, 1), 1); v != 1 { t.Fatalf("SubI: %d", v) } if v, _ := IntAt(MulI(a, 3), 1); v != 6 { t.Fatalf("MulI: %d", v) } f := mustFromFloats(t, []float64{1.0, 2.0}, 2) if v, _ := FloatAt(AddI(f, 5), 0); v != 6.0 { t.Fatalf("AddI on float: %v", v) } if v, _ := FloatAt(AddF(a, 0.5), 0); v != 1.5 { t.Fatalf("AddF: %v", v) } if v, _ := FloatAt(SubF(a, 0.5), 1); v != 1.5 { t.Fatalf("SubF: %v", v) } if v, _ := FloatAt(MulF(a, 1.5), 1); v != 3.0 { t.Fatalf("MulF: %v", v) } // Scalar division is true division regardless of the scalar flavour. div := DivI(a, 2) if div.Dtype() != Float || floatsOf(t, div)[0] != 0.5 { t.Fatalf("DivI: %s %v", div.Dtype(), floatsOf(t, div)) } if v, _ := FloatAt(DivF(a, 4), 1); v != 0.5 { t.Fatalf("DivF: %v", floatsOf(t, DivF(a, 4))) } q, err := QuoI(a, 2) if err != nil { t.Fatalf("QuoI: %v", err) } if v, _ := IntAt(q, 0); v != 0 { t.Fatalf("QuoI: %d", v) } if _, err := QuoI(a, 0); err == nil || !strings.Contains(err.Error(), "division by zero") { t.Fatalf("QuoI zero: %v", err) } if _, err := QuoI(f, 2); err == nil || !strings.Contains(err.Error(), "needs an int array") { t.Fatalf("QuoI float: %v", err) } } // narrowBools builds a bool array for the narrow-dtype contract tests. func narrowBools(t *testing.T, vals []bool, shape ...int) *Array { t.Helper() a, err := FromBools(vals, shape...) if err != nil { t.Fatal(err) } return a } // narrowInt8s builds an int8 array for the narrow-dtype contract tests. func narrowInt8s(t *testing.T, vals []int8, shape ...int) *Array { t.Helper() a, err := FromInt8s(vals, shape...) if err != nil { t.Fatal(err) } return a } // narrowUint8s builds a uint8 array for the narrow-dtype contract tests. func narrowUint8s(t *testing.T, vals []uint8, shape ...int) *Array { t.Helper() a, err := FromUint8s(vals, shape...) if err != nil { t.Fatal(err) } return a } // narrowInt16s builds an int16 array for the narrow-dtype contract tests. func narrowInt16s(t *testing.T, vals []int16, shape ...int) *Array { t.Helper() a, err := FromInt16s(vals, shape...) if err != nil { t.Fatal(err) } return a } // narrowUint32s builds a uint32 array for the narrow-dtype contract tests. func narrowUint32s(t *testing.T, vals []uint32, shape ...int) *Array { t.Helper() a, err := FromUint32s(vals, shape...) if err != nil { t.Fatal(err) } return a } // TestBoolArithmeticIsALoudError pins the contract that bool arrays // carry no arithmetic: every element-wise arithmetic entry whose // promoted dtype is bool answers the named refusal, while comparisons // over bool keep the Int mask. func TestBoolArithmeticIsALoudError(t *testing.T) { x := narrowBools(t, []bool{true, true, false}, 3) y := narrowBools(t, []bool{true, false, false}, 3) entries := []struct { name string fn func() error }{ {"Add", func() error { _, err := Add(x, y); return err }}, {"Sub", func() error { _, err := Sub(x, y); return err }}, {"Mul", func() error { _, err := Mul(x, y); return err }}, {"Div", func() error { _, err := Div(x, y); return err }}, {"Pow", func() error { _, err := Pow(x, y); return err }}, {"Minimum", func() error { _, err := Minimum(x, y); return err }}, {"Maximum", func() error { _, err := Maximum(x, y); return err }}, {"Dot", func() error { _, err := Dot(x, y); return err }}, } for _, tc := range entries { err := tc.fn() if err == nil || !strings.Contains(err.Error(), "bool arrays have no arithmetic") { t.Errorf("%s on bool arrays: %v, want the named arithmetic refusal", tc.name, err) } } // Comparisons are not arithmetic: the mask answers bool, true where // the relation holds. eq, err := Eq(x, y) if err != nil { t.Fatalf("Eq on bool: %v", err) } if eq.Dtype() != Bool { t.Fatalf("Eq on bool answered dtype %s, want the bool mask", eq.Dtype()) } if want := []bool{true, false, true}; !slices.Equal(eq.RawBools(), want) { t.Fatalf("Eq on bool = %v, want %v", eq.RawBools(), want) } // A mixed bool pair promotes into the other operand's dtype, where // arithmetic is defined on the widened values. i8 := narrowInt8s(t, []int8{2, 3, 4}, 3) sum, err := Add(x, i8) if err != nil { t.Fatalf("Add bool with int8: %v", err) } if sum.Dtype() != Int8 { t.Fatalf("Add bool with int8 answered %s, want int8", sum.Dtype()) } if want := []int8{3, 4, 4}; !slices.Equal(sum.RawInt8s(), want) { t.Fatalf("Add bool with int8 = %v, want %v", sum.RawInt8s(), want) } } // TestNarrowIntegerElementwiseContract pins the narrow element-type // loops: same-width dense pairs wrap natively in their own type, mixed // signedness promotes through the containment table, integer true // division answers float64 exactly as the int path does, and the scalar // maps keep their kind with the implicit-store cast. func TestNarrowIntegerElementwiseContract(t *testing.T) { a8 := narrowInt8s(t, []int8{100, -128, 5}, 3) b8 := narrowInt8s(t, []int8{100, 2, 3}, 3) s, err := Add(a8, b8) if err != nil { t.Fatalf("Add int8: %v", err) } if s.Dtype() != Int8 { t.Fatalf("Add int8 answered %s, want int8", s.Dtype()) } if want := []int8{-56, -126, 8}; !slices.Equal(s.RawInt8s(), want) { t.Fatalf("Add int8 wrap = %v, want %v", s.RawInt8s(), want) } m, err := Mul(a8, b8) if err != nil { t.Fatalf("Mul int8: %v", err) } if want := []int8{16, 0, 15}; !slices.Equal(m.RawInt8s(), want) { t.Fatalf("Mul int8 wrap = %v, want %v", m.RawInt8s(), want) } mx, err := Maximum(a8, b8) if err != nil { t.Fatalf("Maximum int8: %v", err) } if want := []int8{100, 2, 5}; !slices.Equal(mx.RawInt8s(), want) { t.Fatalf("Maximum int8 = %v, want %v", mx.RawInt8s(), want) } // Mixed signedness: int8 with uint8 promotes to int16, where both // value sets fit. u8 := narrowUint8s(t, []uint8{200, 200, 200}, 3) mix, err := Add(a8, u8) if err != nil { t.Fatalf("Add int8 with uint8: %v", err) } if mix.Dtype() != Int16 { t.Fatalf("Add int8 with uint8 answered %s, want int16", mix.Dtype()) } if want := []int16{300, 72, 205}; !slices.Equal(mix.RawInt16s(), want) { t.Fatalf("Add int8 with uint8 = %v, want %v", mix.RawInt16s(), want) } // True division: an integer-class pair answers float64, the route // the int pair has always taken. d, err := Div(a8, b8) if err != nil { t.Fatalf("Div int8: %v", err) } if d.Dtype() != Float { t.Fatalf("Div int8 answered %s, want float", d.Dtype()) } got := d.RawFloats() if got[0] != 1 || got[1] != -64 || math.Abs(got[2]-5.0/3.0) > 1e-12 { t.Fatalf("Div int8 = %v, want [1 -64 1.666...]", got) } // Pow keeps the int contract's negative-exponent refusal on the // narrow widths. neg := narrowInt8s(t, []int8{-1}, 1) pos := narrowInt8s(t, []int8{2}, 1) if _, err := Pow(pos, neg); err == nil || !strings.Contains(err.Error(), "negative exponent") { t.Fatalf("Pow int8 with a negative exponent: %v", err) } // The mixed pair whose promote() is Int: the exponent scan must // refuse the negative narrow exponent instead of letting powInt's // loop answer a silent 1. pbase, perr := FromInts([]int64{2, 3}, 2) if perr != nil { t.Fatalf("FromInts: %v", perr) } if _, err := Pow(pbase, narrowInt8s(t, []int8{-3, 4}, 2)); err == nil || !strings.Contains(err.Error(), "negative exponent") { t.Fatalf("Pow int with a negative int8 exponent: %v", err) } if _, err := Pow(narrowUint32s(t, []uint32{2, 3}, 2), narrowInt16s(t, []int16{-1, 2}, 2)); err == nil || !strings.Contains(err.Error(), "negative exponent") { t.Fatalf("Pow uint32 with a negative int16 exponent: %v", err) } // The scalar maps: AddI keeps the kind with the implicit-store cast, // AddF widens the whole integer class to float64, and AddI on a bool // array has no error channel, so it answers nil. si := AddI(a8, 200) if si == nil || si.Dtype() != Int8 { t.Fatalf("AddI int8: %v %v", si, err) } // The int64 sum narrows on store: int8(300) = 44, -128+200 = 72, // int8(205) = -51. if want := []int8{44, 72, -51}; !slices.Equal(si.RawInt8s(), want) { t.Fatalf("AddI int8 = %v, want %v", si.RawInt8s(), want) } sf := AddF(a8, 0.5) if sf.Dtype() != Float { t.Fatalf("AddF int8 answered %s, want float", sf.Dtype()) } if want := []float64{100.5, -127.5, 5.5}; !slices.Equal(sf.RawFloats(), want) { t.Fatalf("AddF int8 = %v, want %v", sf.RawFloats(), want) } bl := narrowBools(t, []bool{true, false}, 2) if got := AddI(bl, 1); got != nil { t.Fatalf("AddI on a bool array = %v, want the nil refusal", got) } }