// Copyright (c) 2026 Petr BalvĂ­n (https://petrbalvin.org) // SPDX-License-Identifier: MIT package spmd import ( "bytes" "math" "testing" "sourcedock.dev/petrbalvin/tensor/internal/core" ) // The wire form is judged by one rule: the array that comes out carries // the exact bits of the array that went in, shape and dtype included. func TestWireRoundTrip(t *testing.T) { nanLow := math.Float64frombits(0x7ff8000000000001) // NaN, low payload bit set nanHigh := math.Float64frombits(0xfff8deadbeef0000) negZero := math.Copysign(0, -1) cases := []*core.Array{ mk(core.FromFloats([]float64{1, -2.5, 0, negZero, math.Inf(1), math.Inf(-1), nanLow, nanHigh, math.MaxFloat64, math.SmallestNonzeroFloat64}, 10)), mk(core.FromFloat32s([]float32{1.5, -0.25, float32(negZero), float32(math.Inf(1)), float32(nanLow)}, 5)), mk(core.FromComplexes([]complex128{1 + 2i, complex(negZero, nanHigh), complex(0, negZero)}, 3)), mk(core.FromInts([]int64{math.MaxInt64, math.MinInt64, -1, 0, 42}, 5)), mk(core.FromBools([]bool{true, false, true, true}, 4)), mk(core.FromInt8s([]int8{math.MinInt8, math.MaxInt8, -1, 0}, 4)), mk(core.FromUint8s([]uint8{0, 255, 128}, 3)), mk(core.FromInt16s([]int16{math.MinInt16, math.MaxInt16, -1}, 3)), mk(core.FromUint16s([]uint16{0, 65535, 32768}, 3)), mk(core.FromInt32s([]int32{math.MinInt32, math.MaxInt32, -1}, 3)), mk(core.FromUint32s([]uint32{0, 4294967295, 2147483648}, 3)), // Float16 raw halves: denormals, infinities, NaN payloads, the // whole span the payload can hold. mk(core.HalvesFromArray([]uint16{0x0001, 0x03ff, 0x7bff, 0x7c00, 0xfc00, 0x7e00, 0x7eaa}, 7)), // Empty and multi-dimensional shapes. mk(core.FromFloats(nil, 0)), mk(core.FromInts([]int64{1, 2, 3, 4, 5, 6}, 2, 3)), mk(core.FromFloats([]float64{1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12}, 2, 1, 2, 3)), } for i, a := range cases { wire, err := encodeArray(nil, a) if err != nil { t.Fatalf("case %d (%s): %v", i, a.Dtype(), err) } b, err := decodeWire(wire) if err != nil { t.Fatalf("case %d (%s): %v", i, a.Dtype(), err) } if b.Dtype() != a.Dtype() { t.Fatalf("case %d: dtype %s came back as %s", i, a.Dtype(), b.Dtype()) } if !sameShape(a.Shape(), b.Shape()) { t.Fatalf("case %d: shape %v came back as %v", i, a.Shape(), b.Shape()) } if !sameBits(a, b) { t.Fatalf("case %d (%s): the round trip moved bits", i, a.Dtype()) } } } // TestWireRefusesTheHostile is the hostile-reader rule: a form that // does not carry exactly what its head names is an error, never a short // read and never a partial answer. func TestWireRefusesTheHostile(t *testing.T) { a := mk(core.FromFloats([]float64{1, 2, 3, 4}, 2, 2)) wire, err := encodeArray(nil, a) if err != nil { t.Fatal(err) } payload := wire[2+8*2:] if _, err := decodeWire(wire[:len(wire)-1]); err == nil { t.Fatal("a short payload decoded") } if _, err := decodeWire(append(bytes.Clone(wire), 0)); err == nil { t.Fatal("a long payload decoded") } if _, err := decodeWire(wire[:1]); err == nil { t.Fatal("a headless form decoded") } if _, err := decodeWire(wire[:6]); err == nil { t.Fatal("a truncated head decoded") } big := []int{math.MaxInt32, math.MaxInt32} if _, err := decodeArray(core.Float, big, nil); err == nil { t.Fatal("an overflowing shape decoded") } if _, err := decodeArray(core.Float, []int{2, -2}, payload); err == nil { t.Fatal("a negative extent decoded") } badBool := mk(core.FromBools([]bool{true, false}, 2)) bw, err := encodeArray(nil, badBool) if err != nil { t.Fatal(err) } bp := bytes.Clone(bw) bp[2+8] = 2 if _, err := decodeWire(bp); err == nil { t.Fatal("a bool byte of 2 decoded") } } // mk builds a fixture or panics; the fixtures are package-level // literals, so a panic lands in the test that declared them. func mk(a *core.Array, err error) *core.Array { if err != nil { panic(err) } return a } // sameBits compares two arrays' raw payloads bit for bit, dtype by // dtype. NaN payloads and signed zeros are part of the contract. func sameBits(a, b *core.Array) bool { if a.Len() != b.Len() || a.Dtype() != b.Dtype() { return false } switch a.Dtype() { case core.Int: return slicesEqual(a.RawInts()[:a.Len()], b.RawInts()[:b.Len()]) case core.Float: x, y := a.RawFloats()[:a.Len()], b.RawFloats()[:b.Len()] for i := range x { if math.Float64bits(x[i]) != math.Float64bits(y[i]) { return false } } return true case core.Float32: x, y := a.RawFloat32s()[:a.Len()], b.RawFloat32s()[:b.Len()] for i := range x { if math.Float32bits(x[i]) != math.Float32bits(y[i]) { return false } } return true case core.Float16: return slicesEqual(a.RawHalves()[:a.Len()], b.RawHalves()[:b.Len()]) case core.Complex: x, y := a.RawComplexes()[:a.Len()], b.RawComplexes()[:b.Len()] for i := range x { if math.Float64bits(real(x[i])) != math.Float64bits(real(y[i])) || math.Float64bits(imag(x[i])) != math.Float64bits(imag(y[i])) { return false } } return true case core.Bool: return slicesEqual(a.RawBools()[:a.Len()], b.RawBools()[:b.Len()]) case core.Int8: return slicesEqual(a.RawInt8s()[:a.Len()], b.RawInt8s()[:b.Len()]) case core.Uint8: return slicesEqual(a.RawUint8s()[:a.Len()], b.RawUint8s()[:b.Len()]) case core.Int16: return slicesEqual(a.RawInt16s()[:a.Len()], b.RawInt16s()[:b.Len()]) case core.Uint16: return slicesEqual(a.RawUint16s()[:a.Len()], b.RawUint16s()[:b.Len()]) case core.Int32: return slicesEqual(a.RawInt32s()[:a.Len()], b.RawInt32s()[:b.Len()]) case core.Uint32: return slicesEqual(a.RawUint32s()[:a.Len()], b.RawUint32s()[:b.Len()]) } return false } func slicesEqual[T comparable](a, b []T) bool { if len(a) != len(b) { return false } for i := range a { if a[i] != b[i] { return false } } return true }