// Copyright (c) 2026 Petr BalvĂ­n (https://petrbalvin.org) // SPDX-License-Identifier: MIT package core import ( "math" "strings" "testing" ) func TestLinspace(t *testing.T) { v, err := Linspace(0, 1, 5) if err != nil { t.Fatal(err) } want := []float64{0, 0.25, 0.5, 0.75, 1} for i := range 5 { if g, _ := FloatAt(v, i); math.Abs(g-want[i]) > 1e-12 { t.Errorf("Linspace[%d]: got %v, want %v", i, g, want[i]) } } one, _ := Linspace(5, 9, 1) if v, _ := FloatAt(one, 0); v != 5 { t.Errorf("Linspace n=1: got %v, want 5", v) } empty, _ := Linspace(0, 1, 0) if empty.Len() != 0 { t.Errorf("Linspace n=0: len %d", empty.Len()) } if _, err := Linspace(0, 1, -1); err == nil { t.Error("Linspace: expected error for negative n") } } func TestRepeat(t *testing.T) { a, _ := FromInts([]int64{1, 2, 3}, 3) r, err := Repeat(a, 2, 0) if err != nil { t.Fatal(err) } if !Equal(mustFromInts(t, []int64{1, 1, 2, 2, 3, 3}, 6), r) { t.Errorf("Repeat: %v", r.RawInts()) } m, _ := FromInts([]int64{1, 2, 3, 4}, 2, 2) r2, err := Repeat(m, 2, 0) if err != nil { t.Fatal(err) } // Repeating along dim 0 duplicates whole rows. if !Equal(mustFromInts(t, []int64{1, 2, 1, 2, 3, 4, 3, 4}, 4, 2), r2) { t.Errorf("Repeat dim0: %v", r2.RawInts()) } if _, err := Repeat(m, 2, 5); err == nil { t.Error("Repeat: expected error for out-of-range dim") } } func TestTile(t *testing.T) { a, _ := FromInts([]int64{1, 2, 3}, 3) tw, err := Tile(a, 2) if err != nil { t.Fatal(err) } if !Equal(mustFromInts(t, []int64{1, 2, 3, 1, 2, 3}, 6), tw) { t.Errorf("Tile: %v", tw.RawInts()) } m, _ := FromInts([]int64{1, 2, 3, 4}, 2, 2) t4, err := Tile(m, 2, 2) if err != nil { t.Fatal(err) } want := []int64{1, 2, 1, 2, 3, 4, 3, 4, 1, 2, 1, 2, 3, 4, 3, 4} if !Equal(mustFromInts(t, want, 4, 4), t4) { t.Errorf("Tile 2x2: %v", t4.RawInts()) } if _, err := Tile(a, -1); err == nil { t.Error("Tile: expected error for negative reps") } } func TestFlip(t *testing.T) { a, _ := FromInts([]int64{1, 2, 3, 4}, 2, 2) f, err := Flip(a) if err != nil { t.Fatal(err) } if !Equal(mustFromInts(t, []int64{4, 3, 2, 1}, 2, 2), f) { t.Errorf("Flip all: %v", f.RawInts()) } f0, err := Flip(a, 0) if err != nil { t.Fatal(err) } if !Equal(mustFromInts(t, []int64{3, 4, 1, 2}, 2, 2), f0) { t.Errorf("Flip dim0: %v", f0.RawInts()) } if _, err := Flip(a, 9); err == nil { t.Error("Flip: expected error for out-of-range dim") } } func TestRoll(t *testing.T) { a, _ := FromInts([]int64{1, 2, 3, 4, 5}, 5) r, err := Roll(a, 2, 0) if err != nil { t.Fatal(err) } if !Equal(mustFromInts(t, []int64{4, 5, 1, 2, 3}, 5), r) { t.Errorf("Roll +2: %v", r.RawInts()) } rNeg, err := Roll(a, -1, 0) if err != nil { t.Fatal(err) } if !Equal(mustFromInts(t, []int64{2, 3, 4, 5, 1}, 5), rNeg) { t.Errorf("Roll -1: %v", rNeg.RawInts()) } if _, err := Roll(a, 1, 3); err == nil { t.Error("Roll: expected error for out-of-range dim") } } func TestUnique(t *testing.T) { a, _ := FromInts([]int64{3, 1, 2, 1, 3, 5}, 6) u, err := Unique(a) if err != nil { t.Fatal(err) } if !Equal(mustFromInts(t, []int64{1, 2, 3, 5}, 4), u) { t.Errorf("Unique int: %v", u.RawInts()) } f, _ := FromFloats([]float64{2.5, math.NaN(), 1.5, 2.5, math.NaN()}, 5) uf, err := Unique(f) if err != nil { t.Fatal(err) } // Sorted: 1.5, 2.5, then one NaN. if v, _ := FloatAt(uf, 0); v != 1.5 { t.Errorf("Unique float[0]: %v", v) } if v, _ := FloatAt(uf, 1); v != 2.5 { t.Errorf("Unique float[1]: %v", v) } if v, _ := FloatAt(uf, 2); !math.IsNaN(v) { t.Errorf("Unique float[2]: %v, want NaN", v) } } func TestArgwhere(t *testing.T) { a, _ := FromInts([]int64{0, 5, 0, 7}, 2, 2) aw, err := Argwhere(a) if err != nil { t.Fatal(err) } if aw.Shape()[0] != 2 || aw.Shape()[1] != 2 { t.Fatalf("Argwhere shape: %v", aw.Shape()) } // Nonzeros at (0,1) and (1,1). if v, _ := IntAt(aw, 0, 0); v != 0 { t.Errorf("Argwhere[0,0]: %v", v) } if v, _ := IntAt(aw, 0, 1); v != 1 { t.Errorf("Argwhere[0,1]: %v", v) } if v, _ := IntAt(aw, 1, 1); v != 1 { t.Errorf("Argwhere[1,1]: %v", v) } } func TestAstype(t *testing.T) { i, _ := FromInts([]int64{1, 2, 3}, 3) f, err := Astype(i, Float) if err != nil { t.Fatal(err) } if f.Dtype() != Float || f.RawFloats()[0] != 1 { t.Errorf("Astype int to float: %s %v", f.Dtype(), f.RawFloats()) } back, err := Astype(f, Int) if err != nil { t.Fatal(err) } if !Equal(back, i) { t.Errorf("Astype round-trip: %v", back.RawInts()) } c, _ := FromComplexes([]complex128{1 + 2i, 3 - 4i}, 2) cf, err := Astype(c, Float) if err != nil { t.Fatal(err) } if cf.RawFloats()[0] != 1 || cf.RawFloats()[1] != 3 { t.Errorf("Astype complex to float: %v", cf.RawFloats()) } if _, err := Astype(c, Int); err == nil || !strings.Contains(err.Error(), "narrow") { t.Errorf("Astype complex to int: %v", err) } } // TestAstypeNarrowDtypes pins the conversion matrix the narrow element // types add: exact widening out of them, range-checked narrowing into // them with the loud first-failure error, the zero-test into bool, the // exact 0/1 widening out of bool, and the legacy pairs keeping their // historical cast semantics beside all of it. func TestAstypeNarrowDtypes(t *testing.T) { // Bool is the zero-test target: NaN reads true, and bool widens out // as 0/1, exact into complex. f, _ := FromFloats([]float64{0, 1.5, math.NaN()}, 3) b, err := Astype(f, Bool) if err != nil { t.Fatal(err) } if got := b.RawBools(); got[0] || !got[1] || !got[2] { t.Errorf("Astype float to bool = %v", got) } cx, _ := FromComplexes([]complex128{0, 3i}, 2) cxb, err := Astype(cx, Bool) if err != nil { t.Fatal(err) } if got := cxb.RawBools(); got[0] || !got[1] { t.Errorf("Astype complex to bool = %v", got) } cb, err := Astype(b, Complex) if err != nil { t.Fatal(err) } if got := cb.RawComplexes(); got[0] != 0 || got[1] != 1 || got[2] != 1 { t.Errorf("Astype bool to complex = %v", got) } bi, err := Astype(b, Int) if err != nil { t.Fatal(err) } if got := bi.RawInts(); got[0] != 0 || got[1] != 1 || got[2] != 1 { t.Errorf("Astype bool to int = %v", got) } // Exact widening out of the narrow integers. i8, _ := FromInt8s([]int8{-128, 0, 127}, 3) wInt, err := Astype(i8, Int) if err != nil { t.Fatal(err) } if got := wInt.RawInts(); got[0] != -128 || got[2] != 127 { t.Errorf("Astype int8 to int = %v", got) } wC, err := Astype(i8, Complex) if err != nil { t.Fatal(err) } if got := wC.RawComplexes(); got[0] != complex(-128, 0) || got[2] != complex(127, 0) { t.Errorf("Astype int8 to complex = %v", got) } u32, _ := FromUint32s([]uint32{4294967295}, 1) wF, err := Astype(u32, Float) if err != nil { t.Fatal(err) } if got := wF.RawFloats(); got[0] != 4294967295 { t.Errorf("Astype uint32 to float = %v", got) } wI, err := Astype(u32, Int) if err != nil { t.Fatal(err) } if got := wI.RawInts(); got[0] != 4294967295 { t.Errorf("Astype uint32 to int = %v", got) } // Narrow to narrow: a contained value set widens exactly, a // non-contained one fails at the first offending index. i16, err := Astype(i8, Int16) if err != nil { t.Fatal(err) } if got := i16.RawInt16s(); got[0] != -128 || got[2] != 127 { t.Errorf("Astype int8 to int16 = %v", got) } if _, err := Astype(i8, Uint8); err == nil || !strings.Contains(err.Error(), "Astype: value -128 at index 0 does not fit uint8") { t.Errorf("Astype int8 to uint8: %v", err) } if _, err := Astype(u32, Int32); err == nil || !strings.Contains(err.Error(), "Astype: value 4294967295 at index 0 does not fit int32") { t.Errorf("Astype uint32 to int32: %v", err) } // An int source checks exactly in int64 space, and the reported // failure is the lowest index even when several elements overflow. big, _ := FromInts([]int64{5, 300, -2}, 3) if _, err := Astype(big, Int8); err == nil || !strings.Contains(err.Error(), "Astype: value 300 at index 1 does not fit int8") { t.Errorf("Astype int to int8: %v", err) } two, _ := FromInts([]int64{400, 300}, 2) if _, err := Astype(two, Int8); err == nil || !strings.Contains(err.Error(), "Astype: value 400 at index 0 does not fit int8") { t.Errorf("Astype int to int8 first failure: %v", err) } // A float source must be finite, integral and in range; the error // names the value in the source's own float64 space. ff, _ := FromFloats([]float64{2, -3, 127}, 3) fi8, err := Astype(ff, Int8) if err != nil { t.Fatal(err) } if got := fi8.RawInt8s(); got[0] != 2 || got[1] != -3 || got[2] != 127 { t.Errorf("Astype float to int8 = %v", got) } fb, _ := FromFloats([]float64{1, 2.5}, 2) if _, err := Astype(fb, Int8); err == nil || !strings.Contains(err.Error(), "Astype: value 2.5 at index 1 does not fit int8") { t.Errorf("Astype float to int8: %v", err) } fInf, _ := FromFloats([]float64{math.Inf(1)}, 1) if _, err := Astype(fInf, Uint16); err == nil || !strings.Contains(err.Error(), "Astype: value +Inf at index 0 does not fit uint16") { t.Errorf("Astype +Inf to uint16: %v", err) } fNaN, _ := FromFloats([]float64{math.NaN()}, 1) if _, err := Astype(fNaN, Int8); err == nil || !strings.Contains(err.Error(), "Astype: value NaN at index 0 does not fit int8") { t.Errorf("Astype NaN to int8: %v", err) } // A complex source into a narrow numeric target keeps the // historical loud refusal; only bool reaches it by zero-test. if _, err := Astype(cx, Uint8); err == nil || !strings.Contains(err.Error(), "cannot narrow complex to uint8") { t.Errorf("Astype complex to uint8: %v", err) } // Legacy pairs keep their cast semantics beside the new rules: // float to int truncates, with no range error. f3, _ := FromFloats([]float64{3.9, -1.2}, 2) li, err := Astype(f3, Int) if err != nil { t.Fatal(err) } if got := li.RawInts(); got[0] != 3 || got[1] != -1 { t.Errorf("Astype float to int legacy truncation = %v", got) } // Same dtype copies through cloneArray for a narrow dtype; Equal // still defaults to the complex payload for these dtypes, so the // comparison reads the payload directly. u8, _ := FromUint8s([]uint8{1, 2, 3}, 3) cp, err := Astype(u8, Uint8) if err != nil { t.Fatal(err) } if cp.Dtype() != Uint8 { t.Fatalf("Astype uint8 to uint8 dtype = %s", cp.Dtype()) } if got := cp.RawUint8s(); got[0] != 1 || got[1] != 2 || got[2] != 3 { t.Errorf("Astype uint8 to uint8 = %v", got) } } func TestItem(t *testing.T) { a, _ := FromFloats([]float64{3.25}, 1) v, err := Item(a) if err != nil { t.Fatal(err) } if v != 3.25 { t.Errorf("Item: %v", v) } b, _ := FromFloats([]float64{1, 2}, 2) if _, err := Item(b); err == nil { t.Error("Item: expected error for multi-element array") } } func TestDiag(t *testing.T) { a, _ := FromInts([]int64{1, 2, 3, 4}, 2, 2) d, err := Diag(a) if err != nil { t.Fatal(err) } if !Equal(mustFromInts(t, []int64{1, 4}, 2), d) { t.Errorf("Diag 2-D: %v", d.RawInts()) } v, _ := FromInts([]int64{5, 6}, 2) m, err := Diag(v) if err != nil { t.Fatal(err) } if m.Shape()[0] != 2 || m.Shape()[1] != 2 { t.Fatalf("Diag 1-D shape: %v", m.Shape()) } if val, _ := IntAt(m, 1, 1); val != 6 { t.Errorf("Diag 1-D [1,1]: %v", val) } }