Files
tensor/io/csv_parity_test.go
T

243 lines
8.1 KiB
Go
Raw Normal View History

2026-09-03 10:00:00 +02:00
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
// SPDX-License-Identifier: MIT
package io
// Parity between the hand-rolled CSV tokenizer and the encoding/csv
// implementation it replaced. The legacy reader below is the oracle:
// the differential test walks a corpus of inputs through both and
// demands the same error text, the same shape and the same float bits,
// and the fuzz target hunts for an input where they part ways.
import (
"bufio"
"encoding/csv"
"io"
"math"
"slices"
"strconv"
"strings"
"testing"
"sourcedock.dev/petrbalvin/tensor/internal/base"
"sourcedock.dev/petrbalvin/tensor/internal/core"
)
// loadCSVReaderLegacy is the encoding/csv reader the tokenizer
// replaced, kept verbatim, byte order mark, buffer growth included.
func loadCSVReaderLegacy(r io.Reader, skipHeader bool) (*core.Array, error) {
br := bufio.NewReader(r)
prefix, _ := br.Peek(3)
if len(prefix) == 3 && prefix[0] == 0xEF && prefix[1] == 0xBB && prefix[2] == 0xBF {
_, _, _ = br.ReadRune() // consume the mark
}
cr := csv.NewReader(br)
cr.FieldsPerRecord = -1 // allow variable; validated manually
// The record is consumed before the next Read, so the reader may
// reuse its slice: one row buffer serves the whole file.
cr.ReuseRecord = true
vals := make([]float64, 0)
rows, cols := 0, 0
header := skipHeader
var ragged, badValue error
for {
rec, err := cr.Read()
if err == io.EOF {
break
}
if err != nil {
return nil, base.Errf("LoadCSV: %w", err)
}
if header {
header = false
continue
}
if rows == 0 {
cols = len(rec)
} else if len(rec) != cols {
if ragged == nil {
ragged = base.Errf("LoadCSV: row %d has %d fields, want %d", rows, len(rec), cols)
}
rows++
continue
}
if badValue == nil {
if need := len(vals) + len(rec); need > cap(vals) {
grown := make([]float64, len(vals), max(2*cap(vals), need))
copy(grown, vals)
vals = grown
}
for j, f := range rec {
v, perr := strconv.ParseFloat(f, 64)
if perr != nil {
badValue = base.Errf("LoadCSV: row %d col %d: %w", rows, j, perr)
break
}
vals = append(vals, v)
}
}
rows++
}
if ragged != nil {
return nil, ragged
}
if badValue != nil {
return nil, badValue
}
return core.FromFloats(vals, rows, cols)
}
// csvParityCases collects the inputs whose treatment the tokenizer has
// to match byte for byte: the shapes a numeric table takes, the quote
// grammar, the line-ending conventions and every defect the reader
// reports.
var csvParityCases = []struct {
name string
in string
header bool
}{
{"plain table", "1.5,2.25\n3.125,4\n", false},
{"last row without a newline", "1,2\n3,4", false},
{"single row without a newline", "1,2", false},
{"crlf", "1,2\r\n3,4\r\n", false},
{"mixed line endings", "1,2\r\n3,4\n5,6\r\n", false},
{"blank lines between records", "1,2\n\n3,4\n\n", false},
{"blank line only", "\n\n\n", false},
{"empty input", "", false},
{"empty field", "1,,3\n4,5,6\n", false},
{"trailing empty field", "1,2,\n3,4,5\n", false},
{"single comma", ",", false},
{"spaces stay in the field", " 1 , 2 \n3,4\n", false},
{"lone carriage return inside a line", "1,2\r3,4\n", false},
{"carriage return before eof", "1,2\r", false},
{"carriage return blank line", "1,2\n\r\n3,4\n", false},
{"bom then table", "\xEF\xBB\xBF1.5,2.5\n3,4\n", false},
{"quoted values", "1,\"2\"\n\"3\",4\n", false},
{"quoted header", "\"a\",\"b\"\n1,2\n", true},
{"quoted comma", "\"1,2\",3\n4,5,6\n", false},
{"escaped quote", "\"a\"\"b\",2\n", false},
{"quote run at field end", "\"a\"\"\",2\n", false},
{"empty quoted fields", "\"\",\"\",\"\"\n", false},
{"multiline quoted field", "\"1\n2\",3\n4,5,6\n", false},
{"multiline quoted field with crlf", "\"1\r\n2\",3\n4,5,6\n", false},
{"blank line inside a quoted field", "\"a\n\nb\",1\n", false},
{"quote only field mid table", "1,2\n\"3\",4\n", false},
{"header skip", "a,b\n1,2\n", true},
{"header skip with blank line", "a,b\n\n1,2\n", true},
{"ragged rows", "1,2\n3\n4,5,6\n", false},
{"ragged rows with header", "a,b\n1,2\n3\n", true},
{"two bad values", "1,x\n4,y\n", false},
{"bad value after ragged row", "1,2\n3\n4,x\n", false},
{"bare quote", "a\"b,1\n", false},
{"bare quote in a later field", "1,b\"c,2\n", false},
{"text after a closing quote", "1,\"a\"b,2\n", false},
{"text after a closing quote on a later line", "1\n\"a\"b,2\n", false},
{"unterminated quote at eof", "1,2\n3,\"ab", false},
{"unterminated quote at end of line", "1,2\n3,\"ab\n", false},
{"unterminated quote after a multiline field", "1,\"a\nb", false},
{"unterminated quote in the header", "\"a,b\n1,2\n", true},
{"only a header", "a,b\n", true},
{"only a header without skip", "a,b\n", false},
{"range overflow", "1e309,2\n", false},
{"range underflow", "1,1e-400\n", false},
{"hex float", "0x1p-2,2\n", false},
{"infinity spelling", "Inf,+Inf,-inf\n", false},
{"nan spelling", "NaN,-nan\n", false},
{"utf-8 in a field", "λ,2\n", false},
{"wide utf-8 in a field", "1,𝄞\n", false},
}
// csvParityExtra builds the generated cases the fixed list cannot
// spell: records longer than the read buffer, quoted or not, and a
// generated table in the shape SaveCSV emits.
func csvParityExtra(t *testing.T) []struct {
name string
in string
header bool
} {
t.Helper()
var sb strings.Builder
for j := range 40000 {
if j > 0 {
sb.WriteByte(',')
}
sb.WriteString(strconv.Itoa(j))
}
longLine := sb.String()
var qb strings.Builder
qb.WriteString("1,\"")
qb.WriteString(strings.Repeat("234567890\n", 10000))
qb.WriteString("\",2\n")
return []struct {
name string
in string
header bool
}{
{"record longer than the read buffer", longLine + "\n1,2\n", false},
{"long record without a newline", longLine, false},
{"long quoted field across the buffer", qb.String(), false},
{"ragged long record", longLine + "\n1\n", false},
{"long last field without a newline", "1," + strings.Repeat("2", 70000), false},
}
}
// checkCSVParity runs one input through both readers and demands the
// same outcome: the same error text, or the same shape and the same
// float bits.
func checkCSVParity(t *testing.T, name, in string, skipHeader bool) {
t.Helper()
want, wantErr := loadCSVReaderLegacy(strings.NewReader(in), skipHeader)
got, gotErr := LoadCSVReader(strings.NewReader(in), skipHeader)
switch {
case wantErr != nil && gotErr != nil:
if wantErr.Error() != gotErr.Error() {
t.Fatalf("%s: error %q, want %q", name, gotErr, wantErr)
}
return
case wantErr != nil || gotErr != nil:
t.Fatalf("%s: error mismatch: legacy %v, tokenizer %v", name, wantErr, gotErr)
}
if !slices.Equal(want.Shape(), got.Shape()) {
t.Fatalf("%s: shape %v, want %v", name, got.Shape(), want.Shape())
}
wv, gv := want.RawFloats(), got.RawFloats()
for i := range wv {
if math.Float64bits(wv[i]) != math.Float64bits(gv[i]) {
t.Fatalf("%s: value %d is %v (%#x), want %v (%#x)", name, i, gv[i], math.Float64bits(gv[i]), wv[i], math.Float64bits(wv[i]))
}
}
}
// TestLoadCSVParity pins the tokenizer to the encoding/csv reader it
// replaced, one input at a time.
func TestLoadCSVParity(t *testing.T) {
for _, tc := range csvParityCases {
checkCSVParity(t, tc.name, tc.in, tc.header)
checkCSVParity(t, tc.name+" with header flag", tc.in, !tc.header)
}
for _, tc := range csvParityExtra(t) {
checkCSVParity(t, tc.name, tc.in, tc.header)
checkCSVParity(t, tc.name+" with header flag", tc.in, !tc.header)
}
// A generated table, the shape the reader is built for.
var sb strings.Builder
a, _ := core.FromFloats([]float64{1.5, -2.25, 3.125, 4, 5.5, -6.75, 1e-9, -0, 0.1, 2, 3, 4}, 4, 3)
if err := SaveCSVWriter(&sb, a); err != nil {
t.Fatal(err)
}
checkCSVParity(t, "saved table", sb.String(), false)
checkCSVParity(t, "saved table with header", "c0,c1,c2\n"+sb.String(), true)
}
// FuzzLoadCSVParity hunts for the input where the tokenizer and the
// encoding/csv reader part ways.
func FuzzLoadCSVParity(f *testing.F) {
for _, tc := range csvParityCases {
f.Add(tc.in, tc.header)
}
f.Add(strings.Repeat("1.5,", 300)+"1.5\n", false)
f.Fuzz(func(t *testing.T, in string, skipHeader bool) {
checkCSVParity(t, "fuzz", in, skipHeader)
})
}