223 lines
7.5 KiB
Go
223 lines
7.5 KiB
Go
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
|
|
// SPDX-License-Identifier: MIT
|
|
|
|
package io
|
|
|
|
import (
|
|
"path/filepath"
|
|
"sourcedock.dev/petrbalvin/tensor/internal/core"
|
|
"strings"
|
|
"testing"
|
|
)
|
|
|
|
func TestCSVRoundTrip(t *testing.T) {
|
|
a, _ := core.FromFloats([]float64{1.5, -2.25, 3.125, 4.0, 5.5, -6.75}, 2, 3)
|
|
dir := t.TempDir()
|
|
path := filepath.Join(dir, "data.csv")
|
|
if err := SaveCSV(path, a); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
back, err := LoadCSV(path, false)
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
if !core.Equal(a, back) {
|
|
t.Errorf("CSV round-trip: got %s, want %s", back, a)
|
|
}
|
|
}
|
|
|
|
func TestCSVWriterReader(t *testing.T) {
|
|
a, _ := core.FromFloats([]float64{1, 2, 3, 4}, 2, 2)
|
|
var sb strings.Builder
|
|
if err := SaveCSVWriter(&sb, a); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
if sb.String() != "1\n2\n3\n4\n" && sb.String() != "1,2\n3,4\n" {
|
|
// csv.Writer separates with commas by default.
|
|
if !strings.Contains(sb.String(), ",") {
|
|
t.Fatalf("unexpected CSV: %q", sb.String())
|
|
}
|
|
}
|
|
back, err := LoadCSVReader(strings.NewReader(sb.String()), false)
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
if !core.Equal(a, back) {
|
|
t.Errorf("reader round-trip: got %s, want %s", back, a)
|
|
}
|
|
}
|
|
|
|
func TestCSVSkipHeader(t *testing.T) {
|
|
data := "a,b,c\n1,2,3\n4,5,6\n"
|
|
back, err := LoadCSVReader(strings.NewReader(data), true)
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
want, _ := core.FromFloats([]float64{1, 2, 3, 4, 5, 6}, 2, 3)
|
|
if !core.Equal(back, want) {
|
|
t.Errorf("skip header: got %s, want %s", back, want)
|
|
}
|
|
}
|
|
|
|
func TestCSVErrors(t *testing.T) {
|
|
// Non-2-D array.
|
|
a, _ := core.FromFloats([]float64{1, 2, 3}, 3)
|
|
if err := SaveCSV("/tmp/x.csv", a); err == nil {
|
|
t.Error("SaveCSV: expected error for 1-D array")
|
|
}
|
|
// Ragged rows.
|
|
if _, err := LoadCSVReader(strings.NewReader("1,2\n3\n"), false); err == nil {
|
|
t.Error("LoadCSV: expected error for ragged rows")
|
|
}
|
|
// Non-numeric.
|
|
if _, err := LoadCSVReader(strings.NewReader("1,x\n"), false); err == nil {
|
|
t.Error("LoadCSV: expected error for non-numeric cell")
|
|
}
|
|
}
|
|
|
|
// TestCSVComplexRefused pins the dtype contract: CSV carries plain
|
|
// float text, so a complex array is refused with an error instead of
|
|
// panicking on its nil float payload.
|
|
func TestCSVComplexRefused(t *testing.T) {
|
|
a, err := core.FromComplexes([]complex128{complex(1, 1), complex(2, 2)}, 2, 1)
|
|
if err != nil {
|
|
t.Fatalf("FromComplexes: %v", err)
|
|
}
|
|
var sb strings.Builder
|
|
if err := SaveCSVWriter(&sb, a); err == nil {
|
|
t.Error("SaveCSVWriter: expected an error for a complex array")
|
|
}
|
|
dir := t.TempDir()
|
|
if err := SaveCSV(filepath.Join(dir, "c.csv"), a); err == nil {
|
|
t.Error("SaveCSV: expected an error for a complex array")
|
|
}
|
|
}
|
|
|
|
// TestLoadCSVErrorPrecedence pins which defect a malformed file
|
|
// reports: a record the csv parser refuses stops the read and is named
|
|
// first, then the first row whose field count differs from the first
|
|
// row's, then the first cell that is not a number, the order a
|
|
// whole-file parse reports them in. The row and cell named are the
|
|
// first offenders, never the last.
|
|
func TestLoadCSVErrorPrecedence(t *testing.T) {
|
|
cases := []struct {
|
|
name string
|
|
in string
|
|
header bool
|
|
want string
|
|
absent string
|
|
}{
|
|
{"a ragged row is named", "1,2\n3\n", false, "row 1 has 1 fields, want 2", ""},
|
|
{"the first ragged row is named", "1,2\n3\n4,5,6\n", false, "row 1 has 1 fields, want 2", "row 2"},
|
|
{"the first bad cell is named", "1,2\n3,x\n4,y\n", false, "row 1 col 1", "row 2"},
|
|
{"a malformed record outranks a ragged row", "1,2\n3\n\"unterminated\n", false, "parse error", "fields, want"},
|
|
{"a ragged row outranks an earlier bad cell", "1,2\nx,2\n3\n", false, "row 2 has 1 fields, want 2", "col"},
|
|
{"the header is not a data row", "a,b\n1,2\n3\n", true, "row 1 has 1 fields, want 2", ""},
|
|
}
|
|
for _, tc := range cases {
|
|
_, err := LoadCSVReader(strings.NewReader(tc.in), tc.header)
|
|
if err == nil {
|
|
t.Fatalf("%s: %q parsed without an error", tc.name, tc.in)
|
|
}
|
|
if !strings.Contains(err.Error(), tc.want) {
|
|
t.Fatalf("%s: %v, want the message to carry %q", tc.name, err, tc.want)
|
|
}
|
|
if tc.absent != "" && strings.Contains(err.Error(), tc.absent) {
|
|
t.Fatalf("%s: %v, want no mention of %q", tc.name, err, tc.absent)
|
|
}
|
|
}
|
|
}
|
|
|
|
// TestCSVWriterDtypes pins the writer's text form for every dtype the
|
|
// core carries: the integer class through exact decimal widenings (an
|
|
// int64 past 2^53 keeps every digit, which a float64 detour would
|
|
// round away), Bool as 0 and 1, the numeric column semantics CSV
|
|
// carries, and Float16 through its exact float64 widening in the same
|
|
// 'g' form the float paths use.
|
|
func TestCSVWriterDtypes(t *testing.T) {
|
|
text := func(t *testing.T, a *core.Array) string {
|
|
t.Helper()
|
|
var sb strings.Builder
|
|
if err := SaveCSVWriter(&sb, a); err != nil {
|
|
t.Fatalf("SaveCSVWriter(%s): %v", a.Dtype(), err)
|
|
}
|
|
return sb.String()
|
|
}
|
|
bools, err := core.FromBools([]bool{true, false, false, true}, 2, 2)
|
|
if err != nil {
|
|
t.Fatalf("FromBools: %v", err)
|
|
}
|
|
if got, want := text(t, bools), "1,0\n0,1\n"; got != want {
|
|
t.Fatalf("bool wrote %q, want %q", got, want)
|
|
}
|
|
i8, err := core.FromInt8s([]int8{-128, 127, 1, -1}, 2, 2)
|
|
if err != nil {
|
|
t.Fatalf("FromInt8s: %v", err)
|
|
}
|
|
if got, want := text(t, i8), "-128,127\n1,-1\n"; got != want {
|
|
t.Fatalf("int8 wrote %q, want %q", got, want)
|
|
}
|
|
u8, err := core.FromUint8s([]uint8{0, 255, 10, 42}, 2, 2)
|
|
if err != nil {
|
|
t.Fatalf("FromUint8s: %v", err)
|
|
}
|
|
if got, want := text(t, u8), "0,255\n10,42\n"; got != want {
|
|
t.Fatalf("uint8 wrote %q, want %q", got, want)
|
|
}
|
|
i16, err := core.FromInt16s([]int16{-32768, 32767, 0, -7}, 2, 2)
|
|
if err != nil {
|
|
t.Fatalf("FromInt16s: %v", err)
|
|
}
|
|
if got, want := text(t, i16), "-32768,32767\n0,-7\n"; got != want {
|
|
t.Fatalf("int16 wrote %q, want %q", got, want)
|
|
}
|
|
u16, err := core.FromUint16s([]uint16{0, 65535, 5, 4096}, 2, 2)
|
|
if err != nil {
|
|
t.Fatalf("FromUint16s: %v", err)
|
|
}
|
|
if got, want := text(t, u16), "0,65535\n5,4096\n"; got != want {
|
|
t.Fatalf("uint16 wrote %q, want %q", got, want)
|
|
}
|
|
i32, err := core.FromInt32s([]int32{-2147483648, 2147483647, 0, -1}, 2, 2)
|
|
if err != nil {
|
|
t.Fatalf("FromInt32s: %v", err)
|
|
}
|
|
if got, want := text(t, i32), "-2147483648,2147483647\n0,-1\n"; got != want {
|
|
t.Fatalf("int32 wrote %q, want %q", got, want)
|
|
}
|
|
u32, err := core.FromUint32s([]uint32{0, 4294967295, 7, 65536}, 2, 2)
|
|
if err != nil {
|
|
t.Fatalf("FromUint32s: %v", err)
|
|
}
|
|
if got, want := text(t, u32), "0,4294967295\n7,65536\n"; got != want {
|
|
t.Fatalf("uint32 wrote %q, want %q", got, want)
|
|
}
|
|
// The %d contract in its element: 2^53+1 cannot survive a float64
|
|
// detour, and the integer path must not take one.
|
|
big, err := core.FromInts([]int64{9007199254740993, -9007199254740993, 0, 2}, 2, 2)
|
|
if err != nil {
|
|
t.Fatalf("FromInts: %v", err)
|
|
}
|
|
if got, want := text(t, big), "9007199254740993,-9007199254740993\n0,2\n"; got != want {
|
|
t.Fatalf("int wrote %q, want %q", got, want)
|
|
}
|
|
half, err := core.FromFloat16s([]float64{0.5, -2, 1, 4}, 2, 2)
|
|
if err != nil {
|
|
t.Fatalf("FromFloat16s: %v", err)
|
|
}
|
|
if got, want := text(t, half), "0.5,-2\n1,4\n"; got != want {
|
|
t.Fatalf("float16 wrote %q, want %q", got, want)
|
|
}
|
|
f32, err := core.FromFloat32s([]float32{1.5, -2.25, 0, 3.25}, 2, 2)
|
|
if err != nil {
|
|
t.Fatalf("FromFloat32s: %v", err)
|
|
}
|
|
if got, want := text(t, f32), "1.5,-2.25\n0,3.25\n"; got != want {
|
|
t.Fatalf("float32 wrote %q, want %q", got, want)
|
|
}
|
|
f64 := mustFloats(t, []float64{1.5, -2.25, 0, 4}, 2, 2)
|
|
if got, want := text(t, f64), "1.5,-2.25\n0,4\n"; got != want {
|
|
t.Fatalf("float wrote %q, want %q", got, want)
|
|
}
|
|
}
|