// Copyright (c) 2026 Petr BalvĂ­n (https://petrbalvin.org) // SPDX-License-Identifier: MIT package io import ( "bufio" "bytes" "encoding/csv" "io" "os" "strconv" "unsafe" "sourcedock.dev/petrbalvin/tensor/internal/base" "sourcedock.dev/petrbalvin/tensor/internal/core" ) // CSV IO. SaveCSV writes a 2-D array as comma-separated text; // LoadCSV reads it back. CSV is the interchange format for tabular // data with spreadsheets and statistical packages, and the plain-text // counterpart to the binary FITS format for data pipelines. Only 2-D // arrays are supported (rows by columns), matching the tabular model. // The writer formats every dtype the core carries: the integer class // as exact integer text, booleans as 0 and 1, floats as float text; // the reader parses everything back as float64. // SaveCSV writes a 2-D array to path as comma-separated values. The // result is loadable by any spreadsheet or data tool. The close error // is part of the write: os.Create's descriptor buffers nothing itself, // but the operating system still reports a failed write through Close // on some filesystems, so a discarded Close error would report a save // that never reached the disk. func SaveCSV(path string, a *core.Array) error { if a.NDim() != 2 { return base.Errf("SaveCSV: needs a 2-D array, got shape %s", base.ShapeText(a.Shape())) } f, err := os.Create(path) if err != nil { return base.Errf("SaveCSV: %w", err) } werr := SaveCSVWriter(f, a) if cerr := f.Close(); werr != nil { return base.Errf("SaveCSV: %w", werr) } else if cerr != nil { return base.Errf("SaveCSV: %w", cerr) } return nil } // SaveCSVWriter writes a 2-D array as CSV to w. Every dtype the core // carries has a text form: Bool as 0 and 1, the integer dtypes as // exact decimal widenings of their payload, Float16 through its exact // float64 widening, and Float32 and Float as 'g' float text. Complex // arrays are refused: CSV carries plain numeric text, and a pair of // raw halves would read back as two unrelated columns. An array with // no columns is refused for the same reason: one empty record per row // is what the writer would emit, and CSV readers treat blank lines as // no records at all, so the shape would come back as (0,0). func SaveCSVWriter(w io.Writer, a *core.Array) error { if a.NDim() != 2 { return base.Errf("SaveCSV: needs a 2-D array, got shape %s", base.ShapeText(a.Shape())) } if a.Dtype() == core.Complex { return base.Errf("SaveCSV: complex arrays are not supported") } cw := csv.NewWriter(w) rows, cols := a.Shape()[0], a.Shape()[1] if cols == 0 { // No column count can follow from CSV text that has no record. return base.Errf("SaveCSV: an array with no columns has no CSV form, got shape %s", base.ShapeText(a.Shape())) } // The package's arrays, views included, always keep payload[i] at // element i (views rebase the payload, never stride it), so the // row walk indexes the payload directly instead of going through // the per-element accessor. rec := make([]string, cols) for r := range rows { baseIdx := r * cols switch a.Dtype() { case core.Float: for c := range cols { rec[c] = strconv.FormatFloat(a.RawFloats()[baseIdx+c], 'g', -1, 64) } case core.Float32: for c := range cols { rec[c] = strconv.FormatFloat(float64(a.RawFloat32s()[baseIdx+c]), 'g', -1, 64) } case core.Float16: // The half payload widens to float64 exactly, and the // widened value formats the way every float path formats. vals := a.RawHalves() for c := range cols { rec[c] = strconv.FormatFloat(core.HalfToFloat64(vals[baseIdx+c]), 'g', -1, 64) } case core.Bool: // Bool writes as 0 and 1: the numeric column semantics a // CSV column carries, the two values the payload stores. vals := a.RawBools() for c := range cols { if vals[baseIdx+c] { rec[c] = "1" } else { rec[c] = "0" } } case core.Int8: vals := a.RawInt8s() for c := range cols { rec[c] = strconv.FormatInt(int64(vals[baseIdx+c]), 10) } case core.Uint8: vals := a.RawUint8s() for c := range cols { rec[c] = strconv.FormatInt(int64(vals[baseIdx+c]), 10) } case core.Int16: vals := a.RawInt16s() for c := range cols { rec[c] = strconv.FormatInt(int64(vals[baseIdx+c]), 10) } case core.Uint16: vals := a.RawUint16s() for c := range cols { rec[c] = strconv.FormatInt(int64(vals[baseIdx+c]), 10) } case core.Int32: vals := a.RawInt32s() for c := range cols { rec[c] = strconv.FormatInt(int64(vals[baseIdx+c]), 10) } case core.Uint32: vals := a.RawUint32s() for c := range cols { rec[c] = strconv.FormatInt(int64(vals[baseIdx+c]), 10) } case core.Int: // The %d contract: an int64 past 2^53 keeps every digit, // where a float64 detour would round it away. vals := a.RawInts() for c := range cols { rec[c] = strconv.FormatInt(vals[baseIdx+c], 10) } default: // Unreachable: Complex is refused above and every other // dtype has its case; the writer answers an error rather // than panic on a payload it cannot name. return base.Errf("SaveCSV: dtype %s has no CSV form", a.Dtype()) } if err := cw.Write(rec); err != nil { return base.Errf("SaveCSV: %w", err) } } cw.Flush() if err := cw.Error(); err != nil { return base.Errf("SaveCSV: %w", err) } return nil } // LoadCSV reads a 2-D float array from a CSV file. Every row must have // the same number of fields; a header row is treated as data unless // skipHeader is true. func LoadCSV(path string, skipHeader bool) (*core.Array, error) { f, err := os.Open(path) if err != nil { return nil, base.Errf("LoadCSV: %w", err) } defer f.Close() return LoadCSVReader(f, skipHeader) } // LoadCSVReader reads a 2-D float array from a CSV stream. A leading // UTF-8 byte order mark is skipped: spreadsheets write one, and the // parser would otherwise glue it onto the first field. // // The records stream: each one is parsed and its values are copied out // as it arrives, so no row of text survives the row that produced it // and the whole file is never held as strings. A malformed record // stops the read and is reported first, then the first row whose field // count differs from the first row's, then the first value that is not // a number, the order a whole-file parse reports them in. // // The record parser is the hand-rolled tokenizer below, tuned for the // numeric tables this entry point serves: it keeps the record // semantics of encoding/csv for this caller's configuration, quotes // and error reports included, while skipping the per-record string the // standard parser builds. func LoadCSVReader(r io.Reader, skipHeader bool) (*core.Array, error) { br := bufio.NewReaderSize(r, csvReadBuffer) prefix, _ := br.Peek(3) if len(prefix) == 3 && prefix[0] == 0xEF && prefix[1] == 0xBB && prefix[2] == 0xBF { _, _, _ = br.ReadRune() // consume the mark } t := &csvTokenizer{br: br} vals := make([]float64, 0) rows, cols := 0, 0 header := skipHeader var ragged, badValue error for { rec, err := t.nextRecord() 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 { // The read carries on to the end even after a defect, so a // later record's syntax error still outranks an earlier // ragged row, as the whole-file parse decides it. The row // counter keeps advancing for the same reason. if ragged == nil { ragged = base.Errf("LoadCSV: row %d has %d fields, want %d", rows, len(rec), cols) } rows++ continue } if badValue == nil { // The buffer doubles as the file arrives: append alone grows a // large float slice by a quarter and rewrites it several times // over, while doubling copies less than the final size once, // and the first record sizes the buffer exactly. 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(csvFieldText(f), 64) if perr != nil { if ne, ok := perr.(*strconv.NumError); ok { // The message carries the field text: rebuild // it over a private copy so it does not hang // off the row buffer this read goes on // overwriting. perr = &strconv.NumError{Func: ne.Func, Num: string(f), Err: ne.Err} } 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) } // csvReadBuffer sizes the read buffer. A numeric table's rows run to a // few hundred bytes, so one refill covers thousands of them and the // reads stop being part of the cost. const csvReadBuffer = 1 << 16 // csvFieldText views a field's bytes as the string strconv.ParseFloat // parses, with no copy. // // SAFETY: the bytes live in the tokenizer's line or record buffer and // stay untouched until the next record is read: ParseFloat reads the // string only within the call, and on failure the NumError is rebuilt // over a private copy before the loop can move on to the next record. // On success no reference to the string escapes. func csvFieldText(f []byte) string { return unsafe.String(unsafe.SliceData(f), len(f)) } // csvTokenizer is the record reader LoadCSVReader runs. It reproduces // what encoding/csv's Reader does for the configuration this package // reads with (comma separator, no comment character, no lazy quotes, // no leading-space trimming, a variable field count): blank lines are // skipped, a quote opens a quoted field, "" escapes one quote, \r\n is // normalised to \n everywhere, interior newlines of a quoted field // included, the last record may lack its newline, a trailing \r is // dropped before EOF, and a malformed record comes back as a // *csv.ParseError naming the same line and column the standard parser // names. What it does not do is build the per-record string: the // fields are byte ranges over the tokenizer's own buffers, valid until // the next call. type csvTokenizer struct { br *bufio.Reader numLine int // the line the reader sits on, counted from one lineBuf []byte // assembles a record longer than the read buffer recordBuf []byte // the record's unescaped fields, back to back fields [][]byte // the current record's fields, reused } // lengthNL reports the number of bytes for the trailing \n. func lengthNL(b []byte) int { if len(b) > 0 && b[len(b)-1] == '\n' { return 1 } return 0 } // readLine returns the next line with its end mark. A trailing \r\n is // normalised to \n in place, a trailing \r is dropped before EOF, and // every line read counts toward the line numbers the parse errors // carry. The result is only valid until the next call. func (t *csvTokenizer) readLine() ([]byte, error) { line, err := t.br.ReadSlice('\n') if err == bufio.ErrBufferFull { t.lineBuf = append(t.lineBuf[:0], line...) for err == bufio.ErrBufferFull { line, err = t.br.ReadSlice('\n') t.lineBuf = append(t.lineBuf, line...) } line = t.lineBuf } readSize := len(line) if readSize > 0 && err == io.EOF { err = nil // For backwards compatibility, drop a trailing \r before EOF. if line[readSize-1] == '\r' { line = line[:readSize-1] } } t.numLine++ // Normalise \r\n to \n on every input line. if n := len(line); n >= 2 && line[n-2] == '\r' && line[n-1] == '\n' { line[n-2] = '\n' line = line[:n-1] } return line, err } // nextRecord parses the next record. The returned slices share the // tokenizer's buffers and are only valid until the next call. func (t *csvTokenizer) nextRecord() ([][]byte, error) { // Read the record's first line, skipping the blank ones. var line []byte var errRead error for errRead == nil { line, errRead = t.readLine() if errRead == nil && len(line) == lengthNL(line) { continue // a blank line carries no record } break } if errRead == io.EOF { return nil, errRead } // Fast path: a record with no quote anywhere splits on the commas // of its one line. A quote is the only construct that carries a // record across lines, escapes a field or raises a parse error, // so the quoted path below owns every case beyond the split. if bytes.IndexByte(line, '"') < 0 { t.fields = t.fields[:0] for { i := bytes.IndexByte(line, ',') if i < 0 { t.fields = append(t.fields, line[:len(line)-lengthNL(line)]) return t.fields, errRead } t.fields = append(t.fields, line[:i]) line = line[i+1:] } } recLine := t.numLine // the line the record starts on t.recordBuf = t.recordBuf[:0] t.fields = t.fields[:0] posLine, posCol := t.numLine, 1 var parseErr error parseField: for { if len(line) == 0 || line[0] != '"' { // Non-quoted field: everything up to the comma, with // the end mark stripped. i := bytes.IndexByte(line, ',') field := line if i >= 0 { field = field[:i] } else { field = field[:len(field)-lengthNL(field)] } // A quote may not appear in a non-quoted field. if j := bytes.IndexByte(field, '"'); j >= 0 { parseErr = &csv.ParseError{StartLine: recLine, Line: t.numLine, Column: posCol + j, Err: csv.ErrBareQuote} break parseField } start := len(t.recordBuf) t.recordBuf = append(t.recordBuf, field...) t.fields = append(t.fields, t.recordBuf[start:len(t.recordBuf)]) if i >= 0 { line = line[i+1:] posCol += i + 1 continue parseField } break parseField } // Quoted field: the opening quote is consumed, the rest // accumulates until the closing one. fieldStart := len(t.recordBuf) line = line[1:] posCol++ for { i := bytes.IndexByte(line, '"') if i >= 0 { t.recordBuf = append(t.recordBuf, line[:i]...) line = line[i+1:] posCol += i + 1 switch { case len(line) > 0 && line[0] == '"': // "" escapes one quote. t.recordBuf = append(t.recordBuf, '"') line = line[1:] posCol++ case len(line) > 0 && line[0] == ',': // ", closes the field. line = line[1:] posCol++ t.fields = append(t.fields, t.recordBuf[fieldStart:len(t.recordBuf)]) continue parseField case len(line) == 0 || (len(line) == 1 && line[0] == '\n'): // A closing quote at the end of the line closes // the field and the record with it. t.fields = append(t.fields, t.recordBuf[fieldStart:len(t.recordBuf)]) break parseField default: // Anything after a closing quote that is neither // comma nor end of line. parseErr = &csv.ParseError{StartLine: recLine, Line: t.numLine, Column: posCol - 1, Err: csv.ErrQuote} break parseField } } else if len(line) > 0 { // End of line inside the field: the whole line, end // mark included, belongs to the field. t.recordBuf = append(t.recordBuf, line...) if errRead != nil { break parseField } posCol += len(line) line, errRead = t.readLine() if len(line) > 0 { posLine++ posCol = 1 } if errRead == io.EOF { errRead = nil } } else { // End of input inside the field. if errRead == nil { parseErr = &csv.ParseError{StartLine: recLine, Line: posLine, Column: posCol, Err: csv.ErrQuote} break parseField } t.fields = append(t.fields, t.recordBuf[fieldStart:len(t.recordBuf)]) break parseField } } } if parseErr == nil { parseErr = errRead } return t.fields, parseErr }