103 lines
3.0 KiB
Go
103 lines
3.0 KiB
Go
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
|
|
// SPDX-License-Identifier: MIT
|
|
|
|
package io
|
|
|
|
import (
|
|
"encoding/hex"
|
|
"strings"
|
|
"testing"
|
|
|
|
"sourcedock.dev/petrbalvin/tensor/internal/core"
|
|
)
|
|
|
|
// TestNetCDFAppendPayloadRefusesNarrowDtypes pins the defensive check
|
|
// in ncAppendPayload: only the dtypes the classic model stores encode,
|
|
// and a narrower array is a named error rather than a read past a
|
|
// payload it does not carry. SaveNetCDF's own validation refuses the
|
|
// same dtypes before the append runs, so the error is unreachable
|
|
// through the public API; the pin holds the direct contract.
|
|
func TestNetCDFAppendPayloadRefusesNarrowDtypes(t *testing.T) {
|
|
f16, err := core.FromFloat16s([]float64{1, 0}, 2)
|
|
if err != nil {
|
|
t.Fatalf("FromFloat16s: %v", err)
|
|
}
|
|
i8, err := core.FromInt8s([]int8{-1, 1}, 2)
|
|
if err != nil {
|
|
t.Fatalf("FromInt8s: %v", err)
|
|
}
|
|
b, err := core.FromBools([]bool{true, false}, 2)
|
|
if err != nil {
|
|
t.Fatalf("FromBools: %v", err)
|
|
}
|
|
for _, tc := range []struct {
|
|
name string
|
|
arr *core.Array
|
|
}{
|
|
{"float16", f16},
|
|
{"int8", i8},
|
|
{"bool", b},
|
|
} {
|
|
t.Run(tc.name, func(t *testing.T) {
|
|
buf, err := ncAppendPayload(nil, tc.arr, 0, tc.arr.Len())
|
|
if err == nil {
|
|
t.Fatalf("ncAppendPayload accepted a %s array", tc.arr.Dtype())
|
|
}
|
|
if buf != nil {
|
|
t.Fatalf("ncAppendPayload returned %d bytes alongside the error", len(buf))
|
|
}
|
|
want := "cannot store dtype " + tc.arr.Dtype().String()
|
|
if !strings.Contains(err.Error(), want) {
|
|
t.Fatalf("error %q does not name the dtype, want %q", err, want)
|
|
}
|
|
})
|
|
}
|
|
}
|
|
|
|
// TestNetCDFAppendPayloadBytePins pins the external encodings byte for
|
|
// byte: float64 as NC_DOUBLE, float32 as NC_FLOAT and int64 as NC_INT,
|
|
// all big-endian, exactly as the classic model stores them.
|
|
func TestNetCDFAppendPayloadBytePins(t *testing.T) {
|
|
f64, err := core.FromFloats([]float64{1.5, -2.25}, 2)
|
|
if err != nil {
|
|
t.Fatalf("FromFloats: %v", err)
|
|
}
|
|
f32, err := core.FromFloat32s([]float32{1.5, -2.25}, 2)
|
|
if err != nil {
|
|
t.Fatalf("FromFloat32s: %v", err)
|
|
}
|
|
ints, err := core.FromInts([]int64{1, -2}, 2)
|
|
if err != nil {
|
|
t.Fatalf("FromInts: %v", err)
|
|
}
|
|
for _, tc := range []struct {
|
|
name string
|
|
arr *core.Array
|
|
want string
|
|
}{
|
|
{"float64", f64, "3ff8000000000000c002000000000000"},
|
|
{"float32", f32, "3fc00000c0100000"},
|
|
{"int64", ints, "00000001fffffffe"},
|
|
} {
|
|
t.Run(tc.name, func(t *testing.T) {
|
|
buf, err := ncAppendPayload(nil, tc.arr, 0, tc.arr.Len())
|
|
if err != nil {
|
|
t.Fatalf("ncAppendPayload: %v", err)
|
|
}
|
|
if got := hex.EncodeToString(buf); got != tc.want {
|
|
t.Fatalf("payload = %s, want %s", got, tc.want)
|
|
}
|
|
})
|
|
}
|
|
// The window form encodes exactly the elements [start, end): a
|
|
// sliced append carries the same bytes the whole array of that
|
|
// window would.
|
|
buf, err := ncAppendPayload(nil, f64, 1, 2)
|
|
if err != nil {
|
|
t.Fatalf("ncAppendPayload: %v", err)
|
|
}
|
|
if got, want := hex.EncodeToString(buf), "c002000000000000"; got != want {
|
|
t.Fatalf("window payload = %s, want %s", got, want)
|
|
}
|
|
}
|