// Copyright (c) 2026 Petr BalvĂ­n (https://petrbalvin.org) // SPDX-License-Identifier: MIT package core import ( "strings" "testing" ) func mustFromComplexes(t *testing.T, vals []complex128, shape ...int) *Array { t.Helper() a, err := FromComplexes(vals, shape...) if err != nil { t.Fatalf("FromComplexes(%v, %v): %v", vals, shape, err) } return a } func TestComplexConstructors(t *testing.T) { a := mustFromComplexes(t, []complex128{complex(1, 2), complex(3, -4)}, 2) if a.Dtype() != Complex || a.Len() != 2 { t.Fatalf("complex array: %s len %d", a.Dtype(), a.Len()) } if v, err := ComplexAt(a, 0); err != nil || v != complex(1, 2) { t.Fatalf("ComplexAt: %v %v", v, err) } if _, err := IntAt(a, 0); err == nil || !strings.Contains(err.Error(), "not int") { t.Fatalf("IntAt on complex: %v", err) } if _, err := FloatAt(a, 0); err == nil || !strings.Contains(err.Error(), "not float") { t.Fatalf("FloatAt on complex: %v", err) } z, err := Zeros(Complex, 2) if err != nil { t.Fatalf("Zeros complex: %v", err) } if v, _ := ComplexAt(z, 1); v != 0 { t.Fatalf("Zeros complex value: %v", v) } one, _ := Ones(Complex, 2) if v, _ := ComplexAt(one, 0); v != 1 { t.Fatalf("Ones complex value: %v", v) } full, _ := FullC(complex(0.5, -0.5), 2) if v, _ := ComplexAt(full, 1); v != complex(0.5, -0.5) { t.Fatalf("FullC value: %v", v) } id, _ := Identity(Complex, 2) if v, _ := ComplexAt(id, 1, 1); v != 1 { t.Fatalf("Identity complex: %v", v) } if v, _ := ComplexAt(id, 0, 1); v != 0 { t.Fatalf("Identity complex off-diagonal: %v", v) } } func TestComplexPromotion(t *testing.T) { c := mustFromComplexes(t, []complex128{complex(1, 1)}, 1) i := mustFromInts(t, []int64{2}, 1) f := mustFromFloats(t, []float64{0.5}, 1) sum, err := Add(c, i) if err != nil || sum.Dtype() != Complex { t.Fatalf("complex + int: %s %v", sum.Dtype(), err) } if v, _ := ComplexAt(sum, 0); v != complex(3, 1) { t.Fatalf("complex + int value: %v", v) } // float + complex promotes too. sum2, _ := Add(f, c) if sum2.Dtype() != Complex { t.Fatalf("float + complex dtype: %s", sum2.Dtype()) } if v, _ := ComplexAt(sum2, 0); v != complex(1.5, 1) { t.Fatalf("float + complex value: %v", v) } // Division on complex stays complex. q, _ := Div(c, f) if q.Dtype() != Complex { t.Fatalf("complex division dtype: %s", q.Dtype()) } if v, _ := ComplexAt(q, 0); v != complex(1, 1)/complex(0.5, 0) { t.Fatalf("complex division value: %v", v) } // Quo rejects complex. if _, err := Quo(c, c); err == nil || !strings.Contains(err.Error(), "needs int arrays") { t.Fatalf("Quo complex: %v", err) } } func TestComplexScalarMirrors(t *testing.T) { c := mustFromComplexes(t, []complex128{complex(1, 2)}, 1) i := mustFromInts(t, []int64{1}, 1) // Int scalars keep complex complex. if v, _ := ComplexAt(AddI(c, 5), 0); v != complex(6, 2) { t.Fatalf("AddI on complex: %v", v) } if v, _ := ComplexAt(MulI(c, 2), 0); v != complex(2, 4) { t.Fatalf("MulI on complex: %v", v) } if v, _ := ComplexAt(DivF(c, 2), 0); v != complex(0.5, 1) { t.Fatalf("DivF on complex: %v", v) } // Complex scalars promote everything. if v, _ := ComplexAt(AddC(i, complex(0, 1)), 0); v != complex(1, 1) { t.Fatalf("AddC on int: %v", v) } out := MulC(i, complex(0, 2)) if out.Dtype() != Complex { t.Fatalf("MulC dtype: %s", out.Dtype()) } if v, _ := ComplexAt(out, 0); v != complex(0, 2) { t.Fatalf("MulC value: %v", v) } if v, _ := ComplexAt(SubC(c, complex(1, 2)), 0); v != 0 { t.Fatalf("SubC value: %v", v) } if v, _ := ComplexAt(DivC(c, complex(0, 1)), 0); v != complex(2, -1) { t.Fatalf("DivC value: %v", v) } if _, err := QuoI(i, 1); err != nil { t.Fatalf("QuoI still works on int: %v", err) } } func TestComplexSumAndDot(t *testing.T) { c := mustFromComplexes(t, []complex128{complex(1, 1), complex(2, -1)}, 2) s := Sum(c) if !s.IsComplex() || s.Complex() != complex(3, 0) { t.Fatalf("Sum complex: %s", s) } d, err := Dot(c, c) if err != nil { t.Fatalf("Dot complex: %v", err) } // (1+i)(1+i) + (2-i)(2-i) = 2i + 3-4i = 3-2i if !d.IsComplex() || d.Complex() != complex(3, -2) { t.Fatalf("Dot complex: %s", d) } // Mixed Dot promotes to complex. dm, _ := Dot(mustFromInts(t, []int64{1, 1}, 2), c) if !dm.IsComplex() { t.Fatalf("Dot mixed: %s", dm) } } func TestComplexUnorderedErrors(t *testing.T) { c := mustFromComplexes(t, []complex128{complex(1, 1)}, 1) if _, err := Min(c); err == nil || !strings.Contains(err.Error(), "no ordering") { t.Fatalf("Min complex: %v", err) } if _, err := Max(c); err == nil || !strings.Contains(err.Error(), "no ordering") { t.Fatalf("Max complex: %v", err) } if _, err := Mean(c); err == nil || !strings.Contains(err.Error(), "no float mean") { t.Fatalf("Mean complex: %v", err) } if _, err := Lt(c, c); err == nil || !strings.Contains(err.Error(), "no ordering") { t.Fatalf("Lt complex: %v", err) } if _, err := GtI(c, 1); err == nil || !strings.Contains(err.Error(), "no ordering") { t.Fatalf("GtI complex: %v", err) } if _, err := LtF(c, 1); err == nil || !strings.Contains(err.Error(), "no ordering") { t.Fatalf("LtF complex: %v", err) } } func TestComplexEqNe(t *testing.T) { a := mustFromComplexes(t, []complex128{complex(1, 2), complex(3, 4)}, 2) b := mustFromComplexes(t, []complex128{complex(1, 2), complex(4, 4)}, 2) eq, err := Eq(a, b) if err != nil { t.Fatalf("Eq complex: %v", err) } if got := maskOf(t, eq); got[0] != 1 || got[1] != 0 { t.Fatalf("Eq complex: %v", got) } ne, _ := Ne(a, b) if got := maskOf(t, ne); got[0] != 0 || got[1] != 1 { t.Fatalf("Ne complex: %v", got) } // Complex vs float compares in complex space; only the element whose // imaginary part is zero can match. realish := mustFromComplexes(t, []complex128{1, complex(3, 4)}, 2) f := mustFromFloats(t, []float64{1, 4}, 2) eqf, _ := Eq(realish, f) if got := maskOf(t, eqf); got[0] != 1 || got[1] != 0 { t.Fatalf("Eq complex-float: %v", got) } } func TestComplexMatMulAndMachinery(t *testing.T) { m := mustFromComplexes(t, []complex128{complex(0, 1), 0, 0, complex(0, 1)}, 2, 2) v := mustFromComplexes(t, []complex128{1, 1}, 2, 1) out, err := MatMul2D(m, v) if err != nil { t.Fatalf("MatMul complex: %v", err) } if out.Dtype() != Complex || out.Shape()[0] != 2 { t.Fatalf("MatMul complex shape: %s", out) } if w, _ := ComplexAt(out, 0, 0); w != complex(0, 1) { t.Fatalf("MatMul complex value: %v", w) } // Transpose, Slice, Row, Mask and WithComplex all carry the dtype. tt := Transpose(m) if tt.Dtype() != Complex { t.Fatalf("Transpose complex dtype: %s", tt.Dtype()) } s, _ := Slice(m, 0, 0, 1) if s.Dtype() != Complex || s.Shape()[0] != 1 { t.Fatalf("Slice complex: %s", s) } r, _ := Row(m, 1) if v2, _ := ComplexAt(r, 1); v2 != complex(0, 1) { t.Fatalf("Row complex: %v", v2) } upd, err := WithComplex(m, complex(9, 9), 0, 0) if err != nil { t.Fatalf("WithComplex: %v", err) } if v2, _ := ComplexAt(upd, 0, 0); v2 != complex(9, 9) { t.Fatalf("WithComplex value: %v", v2) } if v2, _ := ComplexAt(m, 0, 0); v2 != complex(0, 1) { t.Fatalf("WithComplex receiver mutated: %v", v2) } // Where promotes to complex and Mask selects complex elements. cond := mustFromInts(t, []int64{1, 0}, 2) w, _ := Where(cond, mustFromComplexes(t, []complex128{1, 1}, 2), mustFromComplexes(t, []complex128{0, 0}, 2)) if w.Dtype() != Complex { t.Fatalf("Where complex dtype: %s", w.Dtype()) } if v2, _ := ComplexAt(w, 1); v2 != 0 { t.Fatalf("Where complex value: %v", v2) } mask := cond sel, _ := Select(mustFromComplexes(t, []complex128{complex(5, 5), complex(6, 6)}, 2), mask) if sel.Len() != 1 { t.Fatalf("Mask complex len: %d", sel.Len()) } if v2, _ := ComplexAt(sel, 0); v2 != complex(5, 5) { t.Fatalf("Mask complex value: %v", v2) } // String renders complex elements. if got := mustFromComplexes(t, []complex128{complex(1, 2)}, 1).String(); !strings.Contains(got, "(1+2i)") { t.Fatalf("String complex: %q", got) } }