// Copyright (c) 2026 Petr BalvĂ­n (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) } }