feat: initial release
Assisted-by: GLM 5.3 Flash
This commit is contained in:
@@ -0,0 +1,353 @@
|
||||
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
|
||||
// SPDX-License-Identifier: MIT
|
||||
|
||||
package io
|
||||
|
||||
// Benchmarks for the read and write paths that carry the per-value
|
||||
// work: a table decode, a CSV parse, the HDF5 and NetCDF writers and
|
||||
// readers. Every input is deterministic and written once, before the
|
||||
// measured loop.
|
||||
|
||||
import (
|
||||
"encoding/binary"
|
||||
"math"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strconv"
|
||||
"testing"
|
||||
|
||||
"sourcedock.dev/petrbalvin/tensor/internal/core"
|
||||
)
|
||||
|
||||
// Sizes: several thousand rows for the tables, and a few hundred
|
||||
// thousand values for the array formats, which is large enough for the
|
||||
// per-cell work to dominate the fixed cost of each entry point.
|
||||
const (
|
||||
benchRows = 4000
|
||||
benchCols = 24
|
||||
benchSide = 512
|
||||
)
|
||||
|
||||
// benchFloats builds a deterministic float64 array.
|
||||
func benchFloats(shape ...int) *core.Array {
|
||||
a := core.New(core.Float, shape...)
|
||||
raw := a.RawFloats()
|
||||
for i := range raw {
|
||||
raw[i] = math.Sin(float64(i)*0.03125)*1000 + float64(i%97)*0.5
|
||||
}
|
||||
return a
|
||||
}
|
||||
|
||||
// benchInts builds a deterministic int array.
|
||||
func benchInts(shape ...int) *core.Array {
|
||||
a := core.New(core.Int, shape...)
|
||||
raw := a.RawInts()
|
||||
for i := range raw {
|
||||
raw[i] = int64(i)*7 - 3
|
||||
}
|
||||
return a
|
||||
}
|
||||
|
||||
// benchBinaryTableFile writes a BINTABLE holding one column of every
|
||||
// numeric form the reader decodes plus a character column, and returns
|
||||
// its path. The package's own writer emits a subset of those forms
|
||||
// (B, I and J arrive from other writers), so the table is laid out
|
||||
// here.
|
||||
func benchBinaryTableFile(b testing.TB, rows int) string {
|
||||
b.Helper()
|
||||
forms := []string{"K", "D", "E", "J", "I", "B", "L", "8A"}
|
||||
widths := []int{8, 8, 4, 4, 2, 1, 1, 8}
|
||||
rowBytes := 0
|
||||
for _, w := range widths {
|
||||
rowBytes += w
|
||||
}
|
||||
cards := []string{
|
||||
fitsStringCardRaw("XTENSION", "BINTABLE"),
|
||||
fitsIntCard("BITPIX", 8),
|
||||
fitsIntCard("NAXIS", 2),
|
||||
fitsIntCard("NAXIS1", rowBytes),
|
||||
fitsIntCard("NAXIS2", rows),
|
||||
fitsIntCard("PCOUNT", 0),
|
||||
fitsIntCard("GCOUNT", 1),
|
||||
fitsIntCard("TFIELDS", len(forms)),
|
||||
}
|
||||
for i, form := range forms {
|
||||
n := strconv.Itoa(i + 1)
|
||||
cards = append(cards,
|
||||
fitsStringCardRaw("TTYPE"+n, "COL"+n),
|
||||
fitsStringCardRaw("TFORM"+n, form))
|
||||
}
|
||||
cards = append(cards, fitsEndCard())
|
||||
out := fitsAppendCards(nil, []string{
|
||||
fitsBoolCard("SIMPLE", true),
|
||||
fitsIntCard("BITPIX", 8),
|
||||
fitsIntCard("NAXIS", 0),
|
||||
fitsBoolCard("EXTEND", true),
|
||||
fitsEndCard(),
|
||||
})
|
||||
out = fitsAppendCards(out, cards)
|
||||
body := make([]byte, rows*rowBytes)
|
||||
star := []byte("star ")
|
||||
for r := range rows {
|
||||
p := r * rowBytes
|
||||
binary.BigEndian.PutUint64(body[p:], uint64(1000+r))
|
||||
p += 8
|
||||
binary.BigEndian.PutUint64(body[p:], math.Float64bits(float64(r)*0.25-1))
|
||||
p += 8
|
||||
binary.BigEndian.PutUint32(body[p:], math.Float32bits(float32(r)*0.5))
|
||||
p += 4
|
||||
binary.BigEndian.PutUint32(body[p:], uint32(int32(-r)))
|
||||
p += 4
|
||||
binary.BigEndian.PutUint16(body[p:], uint16(int16(r)))
|
||||
p += 2
|
||||
body[p] = byte(r)
|
||||
p++
|
||||
if r%2 == 0 {
|
||||
body[p] = 'T'
|
||||
} else {
|
||||
body[p] = 'F'
|
||||
}
|
||||
p++
|
||||
copy(body[p:], star)
|
||||
}
|
||||
out = append(out, body...)
|
||||
out = fitsAppendZeroPad(out)
|
||||
path := filepath.Join(b.TempDir(), "binary.fits")
|
||||
if err := os.WriteFile(path, out, 0o644); err != nil {
|
||||
b.Fatal(err)
|
||||
}
|
||||
return path
|
||||
}
|
||||
|
||||
// benchASCIITableFile writes an ASCII table with a character, an
|
||||
// integer and a float column, and returns its path.
|
||||
func benchASCIITableFile(b testing.TB, rows int) string {
|
||||
b.Helper()
|
||||
text := make([]string, rows)
|
||||
ints := core.New(core.Int, rows)
|
||||
floats := core.New(core.Float, rows)
|
||||
for i := range rows {
|
||||
text[i] = "star-" + strconv.Itoa(i)
|
||||
ints.RawInts()[i] = int64(1000 + i)
|
||||
floats.RawFloats()[i] = float64(i)*0.125 - 42
|
||||
}
|
||||
cols := []FITSTableColumn{
|
||||
{Name: "STAR", Form: "12A", Text: text},
|
||||
{Name: "ID", Form: "I10", Data: ints},
|
||||
{Name: "MAG", Form: "D20.12", Data: floats},
|
||||
}
|
||||
path := filepath.Join(b.TempDir(), "ascii.fits")
|
||||
if err := SaveFITSTable(path, true, cols, nil); err != nil {
|
||||
b.Fatal(err)
|
||||
}
|
||||
return path
|
||||
}
|
||||
|
||||
// benchCSVFile writes a float64 matrix as CSV and returns its path.
|
||||
func benchCSVFile(b testing.TB, rows, cols int) string {
|
||||
b.Helper()
|
||||
path := filepath.Join(b.TempDir(), "bench.csv")
|
||||
if err := SaveCSV(path, benchFloats(rows, cols)); err != nil {
|
||||
b.Fatal(err)
|
||||
}
|
||||
return path
|
||||
}
|
||||
|
||||
// benchHDF5File writes a classic HDF5 file of one float64 and one
|
||||
// int64 dataset, optionally through the deflate and shuffle filters,
|
||||
// and returns its path.
|
||||
func benchHDF5File(b testing.TB, filtered bool) string {
|
||||
b.Helper()
|
||||
sets := []HDF5Dataset{
|
||||
{Path: "/field", Values: benchFloats(benchSide, benchSide)},
|
||||
{Path: "/ids", Values: benchInts(benchSide, benchSide)},
|
||||
}
|
||||
attrs := map[string]map[string]string{"/": {"origin": "benchmark"}}
|
||||
opts := HDF5WriteOptions{}
|
||||
if filtered {
|
||||
opts = HDF5WriteOptions{Gzip: 6, Shuffle: true}
|
||||
}
|
||||
path := filepath.Join(b.TempDir(), "bench.h5")
|
||||
if err := SaveHDF5(path, sets, attrs, opts); err != nil {
|
||||
b.Fatal(err)
|
||||
}
|
||||
return path
|
||||
}
|
||||
|
||||
// benchNetCDFFile writes a classic NetCDF file of one float64 variable
|
||||
// in two dimensions and returns its path.
|
||||
func benchNetCDFFile(b testing.TB) string {
|
||||
b.Helper()
|
||||
dims := []NetCDFDim{{Name: "row", Length: benchSide}, {Name: "col", Length: benchSide}}
|
||||
vars := []NetCDFVar{{Name: "field", Dims: []string{"row", "col"}, Values: benchFloats(benchSide, benchSide)}}
|
||||
path := filepath.Join(b.TempDir(), "bench.nc")
|
||||
if err := SaveNetCDF(path, dims, vars, map[string]string{"title": "benchmark"}); err != nil {
|
||||
b.Fatal(err)
|
||||
}
|
||||
return path
|
||||
}
|
||||
|
||||
func BenchmarkLoadFITSTableBinary(b *testing.B) {
|
||||
path := benchBinaryTableFile(b, benchRows)
|
||||
b.ReportAllocs()
|
||||
for b.Loop() {
|
||||
table, err := LoadFITSTable(path)
|
||||
if err != nil {
|
||||
b.Fatal(err)
|
||||
}
|
||||
if table.Rows != benchRows || len(table.Columns) != 8 || table.Text[7] == nil {
|
||||
b.Fatalf("table came back as %s with %d columns", table.Kind, len(table.Columns))
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func BenchmarkLoadFITSTableASCII(b *testing.B) {
|
||||
path := benchASCIITableFile(b, benchRows)
|
||||
b.ReportAllocs()
|
||||
for b.Loop() {
|
||||
table, err := LoadFITSTable(path)
|
||||
if err != nil {
|
||||
b.Fatal(err)
|
||||
}
|
||||
if table.Rows != benchRows || len(table.Columns) != 3 {
|
||||
b.Fatalf("table came back with %d rows and %d columns", table.Rows, len(table.Columns))
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func BenchmarkLoadCSV(b *testing.B) {
|
||||
path := benchCSVFile(b, benchRows, benchCols)
|
||||
b.ReportAllocs()
|
||||
for b.Loop() {
|
||||
a, err := LoadCSV(path, false)
|
||||
if err != nil {
|
||||
b.Fatal(err)
|
||||
}
|
||||
if a.Shape()[0] != benchRows || a.Shape()[1] != benchCols {
|
||||
b.Fatalf("array came back with shape %v", a.Shape())
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func BenchmarkSaveHDF5(b *testing.B) {
|
||||
sets := []HDF5Dataset{
|
||||
{Path: "/field", Values: benchFloats(benchSide, benchSide)},
|
||||
{Path: "/ids", Values: benchInts(benchSide, benchSide)},
|
||||
}
|
||||
attrs := map[string]map[string]string{"/": {"origin": "benchmark"}}
|
||||
path := filepath.Join(b.TempDir(), "bench.h5")
|
||||
b.ReportAllocs()
|
||||
b.SetBytes(int64(2 * benchSide * benchSide * 8))
|
||||
for b.Loop() {
|
||||
if err := SaveHDF5(path, sets, attrs); err != nil {
|
||||
b.Fatal(err)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// BenchmarkSaveHDF5Large writes one 8 MiB float64 field, the size
|
||||
// where moving the payload around shows above the per-value work.
|
||||
func BenchmarkSaveHDF5Large(b *testing.B) {
|
||||
sets := []HDF5Dataset{{Path: "/field", Values: benchFloats(1024, 1024)}}
|
||||
path := filepath.Join(b.TempDir(), "bench.h5")
|
||||
b.ReportAllocs()
|
||||
b.SetBytes(8 << 20)
|
||||
for b.Loop() {
|
||||
if err := SaveHDF5(path, sets, nil); err != nil {
|
||||
b.Fatal(err)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// BenchmarkSaveHDF5Filtered writes the two benchmark datasets through
|
||||
// the shuffle and deflate filters, the chunked path.
|
||||
func BenchmarkSaveHDF5Filtered(b *testing.B) {
|
||||
sets := []HDF5Dataset{
|
||||
{Path: "/field", Values: benchFloats(benchSide, benchSide)},
|
||||
{Path: "/ids", Values: benchInts(benchSide, benchSide)},
|
||||
}
|
||||
path := filepath.Join(b.TempDir(), "bench.h5")
|
||||
b.ReportAllocs()
|
||||
b.SetBytes(int64(2 * benchSide * benchSide * 8))
|
||||
for b.Loop() {
|
||||
if err := SaveHDF5(path, sets, nil, HDF5WriteOptions{Gzip: 6, Shuffle: true}); err != nil {
|
||||
b.Fatal(err)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// BenchmarkSaveHDF5Text writes one string dataset, the fixed-width
|
||||
// text path whose elements are padded to the longest of the file.
|
||||
func BenchmarkSaveHDF5Text(b *testing.B) {
|
||||
const side = 512
|
||||
text := make([]string, side*side)
|
||||
width := 0
|
||||
for i := range text {
|
||||
text[i] = "row" + strconv.Itoa(i)
|
||||
width = max(width, len(text[i]))
|
||||
}
|
||||
sets := []HDF5TextDataset{{Path: "/labels", Shape: []int{side, side}, Text: text}}
|
||||
path := filepath.Join(b.TempDir(), "bench.h5")
|
||||
b.ReportAllocs()
|
||||
b.SetBytes(int64(side * side * width))
|
||||
for b.Loop() {
|
||||
if err := SaveHDF5Text(path, sets); err != nil {
|
||||
b.Fatal(err)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func BenchmarkLoadHDF5(b *testing.B) {
|
||||
path := benchHDF5File(b, false)
|
||||
b.ReportAllocs()
|
||||
for b.Loop() {
|
||||
sets, err := LoadHDF5(path)
|
||||
if err != nil {
|
||||
b.Fatal(err)
|
||||
}
|
||||
if len(sets) != 2 {
|
||||
b.Fatalf("file came back with %d datasets", len(sets))
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func BenchmarkLoadHDF5Filtered(b *testing.B) {
|
||||
path := benchHDF5File(b, true)
|
||||
b.ReportAllocs()
|
||||
for b.Loop() {
|
||||
sets, err := LoadHDF5(path)
|
||||
if err != nil {
|
||||
b.Fatal(err)
|
||||
}
|
||||
if len(sets) != 2 {
|
||||
b.Fatalf("file came back with %d datasets", len(sets))
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func BenchmarkSaveNetCDF(b *testing.B) {
|
||||
dims := []NetCDFDim{{Name: "row", Length: benchSide}, {Name: "col", Length: benchSide}}
|
||||
vars := []NetCDFVar{{Name: "field", Dims: []string{"row", "col"}, Values: benchFloats(benchSide, benchSide)}}
|
||||
attrs := map[string]string{"title": "benchmark"}
|
||||
path := filepath.Join(b.TempDir(), "bench.nc")
|
||||
b.ReportAllocs()
|
||||
for b.Loop() {
|
||||
if err := SaveNetCDF(path, dims, vars, attrs); err != nil {
|
||||
b.Fatal(err)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func BenchmarkLoadNetCDF(b *testing.B) {
|
||||
path := benchNetCDFFile(b)
|
||||
b.ReportAllocs()
|
||||
for b.Loop() {
|
||||
_, vars, _, err := LoadNetCDF(path)
|
||||
if err != nil {
|
||||
b.Fatal(err)
|
||||
}
|
||||
if len(vars) != 1 {
|
||||
b.Fatalf("file came back with %d variables", len(vars))
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,471 @@
|
||||
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (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
|
||||
}
|
||||
@@ -0,0 +1,65 @@
|
||||
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
|
||||
// SPDX-License-Identifier: MIT
|
||||
|
||||
package io
|
||||
|
||||
// The A/B benchmark for the CSV read path: the legacy encoding/csv
|
||||
// reader and the hand-rolled tokenizer side by side, in one process.
|
||||
// The two sub-benchmarks alternate within every -count round, so both
|
||||
// see the same machine, and the decision reads the medians with 1 to 2
|
||||
// percent treated as noise.
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"io"
|
||||
"testing"
|
||||
|
||||
"sourcedock.dev/petrbalvin/tensor/internal/core"
|
||||
)
|
||||
|
||||
// The A/B table: a hundred thousand rows of sixteen float64 values, in
|
||||
// the text shape SaveCSV emits, sized so the per-value work dominates
|
||||
// everything else.
|
||||
const (
|
||||
benchCSVABRows = 100000
|
||||
benchCSVABCols = 16
|
||||
)
|
||||
|
||||
// benchCSVABData renders the A/B table as CSV bytes, once, before the
|
||||
// measured loops.
|
||||
func benchCSVABData(b *testing.B) []byte {
|
||||
b.Helper()
|
||||
var buf bytes.Buffer
|
||||
if err := SaveCSVWriter(&buf, benchFloats(benchCSVABRows, benchCSVABCols)); err != nil {
|
||||
b.Fatal(err)
|
||||
}
|
||||
return buf.Bytes()
|
||||
}
|
||||
|
||||
// BenchmarkLoadCSVReader races the two readers over the same table.
|
||||
func BenchmarkLoadCSVReader(b *testing.B) {
|
||||
data := benchCSVABData(b)
|
||||
impls := []struct {
|
||||
name string
|
||||
load func(io.Reader, bool) (*core.Array, error)
|
||||
}{
|
||||
{"old", loadCSVReaderLegacy},
|
||||
{"new", LoadCSVReader},
|
||||
}
|
||||
for _, impl := range impls {
|
||||
b.Run(impl.name, func(b *testing.B) {
|
||||
r := bytes.NewReader(data)
|
||||
b.ReportAllocs()
|
||||
for b.Loop() {
|
||||
r.Reset(data)
|
||||
a, err := impl.load(r, false)
|
||||
if err != nil {
|
||||
b.Fatal(err)
|
||||
}
|
||||
if a.Shape()[0] != benchCSVABRows || a.Shape()[1] != benchCSVABCols {
|
||||
b.Fatalf("array came back with shape %v", a.Shape())
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,242 @@
|
||||
// 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)
|
||||
})
|
||||
}
|
||||
+222
@@ -0,0 +1,222 @@
|
||||
// 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)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,61 @@
|
||||
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
|
||||
// SPDX-License-Identifier: MIT
|
||||
|
||||
// Package io reads and writes the formats scientific data arrives in:
|
||||
// comma-separated text, FITS images and tables, HDF5, NetCDF classic
|
||||
// and native-endian memory maps. Every loader returns the library's
|
||||
// own Array, so a file read is an ordinary value the rest of the
|
||||
// library operates on, and every writer takes one.
|
||||
//
|
||||
// CSV is the interchange format for tabular data with spreadsheets
|
||||
// and statistical packages: LoadCSV, LoadCSVReader, SaveCSV and
|
||||
// SaveCSVWriter handle 2-D arrays, with or without a header row. The
|
||||
// writer formats every numeric dtype the core carries: the integer
|
||||
// class as exact integer text, booleans as 0 and 1, floats as float
|
||||
// text; complex is refused, because CSV carries plain numeric text
|
||||
// and a pair of raw halves would read back as two unrelated columns.
|
||||
// The readers parse everything back as float64.
|
||||
//
|
||||
// FITS is astronomy's archival format. LoadFITS reads a primary image
|
||||
// (BITPIX -64 and -32) with its header cards and applies the
|
||||
// BSCALE/BZERO scaling; SaveFITS writes one. LoadFITSTable and
|
||||
// SaveFITSTable cover the binary and ASCII table extensions, where
|
||||
// catalogues and observation logs live.
|
||||
//
|
||||
// HDF5 is read by LoadHDF5, which returns every dataset of a file by
|
||||
// path, with the attributes of the groups it sits in merged into it.
|
||||
// Superblocks 0 to 3, object headers of version 1 and 2, symbol-table
|
||||
// and link-message groups, contiguous, compact and chunked storage and
|
||||
// the deflate, shuffle and fletcher32 filters are supported. Dataset
|
||||
// values land the core dtype their datatype declares: fixed-point data
|
||||
// by stored width and signedness, the boolean enumeration convention
|
||||
// as Bool, floating-point data as float64 or float32 by width. What is
|
||||
// not supported, among it dense groups, the version 2 chunk B-tree,
|
||||
// string datasets, every big-endian datatype, bit fields, non-boolean
|
||||
// enumerations and unsigned 64-bit integers, is refused with an error
|
||||
// naming it. SaveHDF5 writes the mirror image in the classic layout
|
||||
// or, with Latest, the superblock 3 layout, and SaveHDF5Text writes
|
||||
// fixed-length string datasets.
|
||||
//
|
||||
// NetCDF classic (CDF-1 and CDF-2) is the archival format of climate
|
||||
// and ocean science: LoadNetCDF returns named dimensions, variables
|
||||
// and global attributes, and SaveNetCDF writes CDF-1. Each variable
|
||||
// lands the core dtype its classic type code carries: NC_BYTE as
|
||||
// int8, NC_CHAR as uint8 raw bytes (CHAR carries bytes at the array
|
||||
// level, never text), NC_SHORT as int16, NC_INT as int32, and
|
||||
// NC_FLOAT and NC_DOUBLE as float64. A variable that lands a narrow
|
||||
// dtype from a file is refused by SaveNetCDF, whose writer stores
|
||||
// float64, float32 and int64 arrays only; convert with Astype first.
|
||||
// Record dimensions are read and written in both directions; the
|
||||
// writer stores float64, float32 and int64 arrays, and a type code
|
||||
// beyond the classic six is refused by name.
|
||||
//
|
||||
// MapFloats, MapFloat32s and MapInts open a native-endian file as a
|
||||
// read-only array without reading it, which suits data far larger than
|
||||
// memory; release unmaps it and comes strictly last. SaveNativeFloats
|
||||
// writes the format MapFloats reads.
|
||||
//
|
||||
// Each format is covered by a documented subset and the rest is
|
||||
// refused by name, never half-read. Errors carry the library's
|
||||
// "tensor: " prefix and the name of the call that produced them.
|
||||
package io
|
||||
@@ -0,0 +1,237 @@
|
||||
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
|
||||
// SPDX-License-Identifier: MIT
|
||||
|
||||
package io_test
|
||||
|
||||
// The godoc examples: one runnable, checked snippet per format the
|
||||
// package speaks. `go test` executes them, so the documentation cannot
|
||||
// rot. Each one writes into a fresh temporary directory and removes it
|
||||
// again.
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"log"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
|
||||
tensor "sourcedock.dev/petrbalvin/tensor"
|
||||
"sourcedock.dev/petrbalvin/tensor/io"
|
||||
)
|
||||
|
||||
// tempDir makes a fresh temporary directory for one example; the
|
||||
// example removes it with a deferred os.RemoveAll.
|
||||
func tempDir() string {
|
||||
dir, err := os.MkdirTemp("", "tensor-io-example-")
|
||||
if err != nil {
|
||||
log.Fatal(err)
|
||||
}
|
||||
return dir
|
||||
}
|
||||
|
||||
// A 2-D array goes out as comma-separated text and reads back as the
|
||||
// same numbers.
|
||||
func ExampleSaveCSV() {
|
||||
dir := tempDir()
|
||||
defer os.RemoveAll(dir)
|
||||
path := filepath.Join(dir, "readings.csv")
|
||||
|
||||
a, err := tensor.FromFloats([]float64{18.5, 21.25, 19.75, 23, 20.5, 17.25}, 2, 3)
|
||||
if err != nil {
|
||||
log.Fatal(err)
|
||||
}
|
||||
if err := io.SaveCSV(path, a); err != nil {
|
||||
log.Fatal(err)
|
||||
}
|
||||
back, err := io.LoadCSV(path, false)
|
||||
if err != nil {
|
||||
log.Fatal(err)
|
||||
}
|
||||
fmt.Println("shape:", back.Shape())
|
||||
fmt.Println("first:", back.FloatAt(0), "last:", back.FloatAt(back.Len()-1))
|
||||
// Output:
|
||||
// shape: [2 3]
|
||||
// first: 18.5 last: 17.25
|
||||
}
|
||||
|
||||
// The stream form writes CSV to any writer and reads it from any
|
||||
// reader, skipping a header row on the way back.
|
||||
func ExampleSaveCSVWriter() {
|
||||
a, err := tensor.FromFloats([]float64{1, 2, 3, 4}, 2, 2)
|
||||
if err != nil {
|
||||
log.Fatal(err)
|
||||
}
|
||||
var buf strings.Builder
|
||||
if err := io.SaveCSVWriter(&buf, a); err != nil {
|
||||
log.Fatal(err)
|
||||
}
|
||||
fmt.Print(buf.String())
|
||||
|
||||
back, err := io.LoadCSVReader(strings.NewReader("row,c1,c2\n"+buf.String()), true)
|
||||
if err != nil {
|
||||
log.Fatal(err)
|
||||
}
|
||||
fmt.Println("shape:", back.Shape(), "last:", back.FloatAt(3))
|
||||
// Output:
|
||||
// 1,2
|
||||
// 3,4
|
||||
// shape: [2 2] last: 4
|
||||
}
|
||||
|
||||
// A float64 image goes out as a FITS primary image with header cards
|
||||
// and reads back with the cards beside the values.
|
||||
func ExampleSaveFITS() {
|
||||
dir := tempDir()
|
||||
defer os.RemoveAll(dir)
|
||||
path := filepath.Join(dir, "image.fits")
|
||||
|
||||
a, err := tensor.FromFloats([]float64{1, 2.5, 3, 4, 5, 6}, 2, 3)
|
||||
if err != nil {
|
||||
log.Fatal(err)
|
||||
}
|
||||
headers := map[string]string{"OBJECT": "M31", "EXPTIME": "600"}
|
||||
if err := io.SaveFITS(path, a, headers); err != nil {
|
||||
log.Fatal(err)
|
||||
}
|
||||
img, header, err := io.LoadFITS(path)
|
||||
if err != nil {
|
||||
log.Fatal(err)
|
||||
}
|
||||
fmt.Println("shape:", img.Shape(), "dtype:", img.Dtype())
|
||||
fmt.Println("OBJECT:", header["OBJECT"], "EXPTIME:", header["EXPTIME"])
|
||||
fmt.Println("last:", img.FloatAt(img.Len()-1))
|
||||
// Output:
|
||||
// shape: [2 3] dtype: float
|
||||
// OBJECT: M31 EXPTIME: 600
|
||||
// last: 6
|
||||
}
|
||||
|
||||
// A table extension holds a character column and a numeric column; the
|
||||
// reader returns both, parallel to the file's column list.
|
||||
func ExampleSaveFITSTable() {
|
||||
dir := tempDir()
|
||||
defer os.RemoveAll(dir)
|
||||
path := filepath.Join(dir, "catalogue.fits")
|
||||
|
||||
flux, err := tensor.FromFloats([]float64{1.5, 2.25, 3.75}, 3)
|
||||
if err != nil {
|
||||
log.Fatal(err)
|
||||
}
|
||||
cols := []io.FITSTableColumn{
|
||||
{Name: "source", Form: "8A", Text: []string{"alpha", "beta", "gamma"}},
|
||||
{Name: "flux", Unit: "Jy", Form: "D", Data: flux},
|
||||
}
|
||||
if err := io.SaveFITSTable(path, false, cols, nil); err != nil {
|
||||
log.Fatal(err)
|
||||
}
|
||||
table, err := io.LoadFITSTable(path)
|
||||
if err != nil {
|
||||
log.Fatal(err)
|
||||
}
|
||||
fmt.Println("kind:", table.Kind, "rows:", table.Rows)
|
||||
fmt.Println("names:", table.Names, "unit:", table.Units[1])
|
||||
fmt.Println("source[1]:", table.Text[0][1], "flux[1]:", table.Columns[1].FloatAt(1))
|
||||
// Output:
|
||||
// kind: BINTABLE rows: 3
|
||||
// names: [source flux] unit: Jy
|
||||
// source[1]: beta flux[1]: 2.25
|
||||
}
|
||||
|
||||
// Two datasets, one in a group, with attributes on the datasets and on
|
||||
// the root and the group: the file reads back with the same paths,
|
||||
// shapes and dtypes, and each dataset carries the attributes of the
|
||||
// enclosing groups.
|
||||
func ExampleSaveHDF5() {
|
||||
dir := tempDir()
|
||||
defer os.RemoveAll(dir)
|
||||
path := filepath.Join(dir, "scan.h5")
|
||||
|
||||
temp, err := tensor.FromFloat32s([]float32{1.5, 2.5, 3.5, 4.5}, 2, 2)
|
||||
if err != nil {
|
||||
log.Fatal(err)
|
||||
}
|
||||
counts, err := tensor.FromInts([]int64{100, 200, 300}, 3)
|
||||
if err != nil {
|
||||
log.Fatal(err)
|
||||
}
|
||||
sets := []io.HDF5Dataset{
|
||||
{Path: "/scan/temperature", Values: temp, Attrs: map[string]string{"units": "degC"}},
|
||||
{Path: "/scan/background", Values: counts, Attrs: map[string]string{"units": "counts"}},
|
||||
}
|
||||
groupAttrs := map[string]map[string]string{
|
||||
"/": {"title": "cruise"},
|
||||
"/scan": {"instrument": "thermistor"},
|
||||
}
|
||||
if err := io.SaveHDF5(path, sets, groupAttrs); err != nil {
|
||||
log.Fatal(err)
|
||||
}
|
||||
got, err := io.LoadHDF5(path)
|
||||
if err != nil {
|
||||
log.Fatal(err)
|
||||
}
|
||||
for _, d := range got {
|
||||
fmt.Printf("%s %v %s %s %s\n", d.Path, d.Shape, d.Values.Dtype(), d.Attrs["instrument"], d.Attrs["units"])
|
||||
}
|
||||
// Output:
|
||||
// /scan/background [3] int thermistor counts
|
||||
// /scan/temperature [2 2] float32 thermistor degC
|
||||
}
|
||||
|
||||
// Dimensions, a variable with an attribute and a global attribute go
|
||||
// through NetCDF classic and back; the writer stores float64 as
|
||||
// NC_DOUBLE, which the reader lands as float64 again, and the integer
|
||||
// type codes land their own dtypes the same way.
|
||||
func ExampleSaveNetCDF() {
|
||||
dir := tempDir()
|
||||
defer os.RemoveAll(dir)
|
||||
path := filepath.Join(dir, "field.nc")
|
||||
|
||||
dims := []io.NetCDFDim{{Name: "lat", Length: 2}, {Name: "lon", Length: 3}}
|
||||
temp, err := tensor.FromFloats([]float64{1, 2, 3, 4, 5, 6}, 2, 3)
|
||||
if err != nil {
|
||||
log.Fatal(err)
|
||||
}
|
||||
vars := []io.NetCDFVar{{
|
||||
Name: "temp",
|
||||
Dims: []string{"lat", "lon"},
|
||||
Values: temp,
|
||||
Attrs: map[string]string{"units": "degC"},
|
||||
}}
|
||||
if err := io.SaveNetCDF(path, dims, vars, map[string]string{"title": "cruise"}); err != nil {
|
||||
log.Fatal(err)
|
||||
}
|
||||
gotDims, gotVars, gotAttrs, err := io.LoadNetCDF(path)
|
||||
if err != nil {
|
||||
log.Fatal(err)
|
||||
}
|
||||
fmt.Println("dims:", gotDims)
|
||||
fmt.Println("var:", gotVars[0].Name, gotVars[0].Dims, gotVars[0].Attrs["units"])
|
||||
fmt.Println("title:", gotAttrs["title"], "last:", gotVars[0].Values.FloatAt(5))
|
||||
// Output:
|
||||
// dims: [{lat 2} {lon 3}]
|
||||
// var: temp [lat lon] degC
|
||||
// title: cruise last: 6
|
||||
}
|
||||
|
||||
// A native-endian file of float64 values maps into a read-only array
|
||||
// without being read; release unmaps it and comes strictly last.
|
||||
func ExampleSaveNativeFloats() {
|
||||
dir := tempDir()
|
||||
defer os.RemoveAll(dir)
|
||||
path := filepath.Join(dir, "values.bin")
|
||||
|
||||
values := []float64{1.5, 2.5, 3.5, 4.5, 5.5}
|
||||
if err := io.SaveNativeFloats(path, values); err != nil {
|
||||
log.Fatal(err)
|
||||
}
|
||||
a, release, err := io.MapFloats(path, 0, len(values))
|
||||
if err != nil {
|
||||
log.Fatal(err)
|
||||
}
|
||||
fmt.Println("mapped:", a.Len(), a.Dtype(), a.FloatAt(0), a.FloatAt(4))
|
||||
if err := release(); err != nil {
|
||||
log.Fatal(err)
|
||||
}
|
||||
// Output:
|
||||
// mapped: 5 float 1.5 5.5
|
||||
}
|
||||
@@ -0,0 +1,573 @@
|
||||
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
|
||||
// SPDX-License-Identifier: MIT
|
||||
|
||||
package io
|
||||
|
||||
import (
|
||||
"encoding/binary"
|
||||
"fmt"
|
||||
"math"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"runtime"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"sourcedock.dev/petrbalvin/tensor/internal/core"
|
||||
)
|
||||
|
||||
// Regression tests for hostile inputs across the io package: the size
|
||||
// arithmetic of the HDF5 reader, the NetCDF record-size
|
||||
// pre-pass, the FITS table encoders and decoders, the mmap count
|
||||
// conversion, and the CSV shape contract. Every hostile file below is
|
||||
// built by hand, byte by byte, because the point of each test is a
|
||||
// declared size the writer of a well-formed file would never produce.
|
||||
|
||||
// h5HostileFile returns an n-byte HDF5 file with the signature and a
|
||||
// version 0 superblock (the classic layout: eight-byte addresses and
|
||||
// lengths) whose root object header sits at offset 96.
|
||||
func h5HostileFile(n int) []byte {
|
||||
f := make([]byte, n)
|
||||
copy(f, hdf5Magic)
|
||||
f[8] = 0 // superblock version 0
|
||||
f[13] = 8
|
||||
f[14] = 8
|
||||
binary.LittleEndian.PutUint64(f[32:], math.MaxUint64) // free space undefined
|
||||
binary.LittleEndian.PutUint64(f[40:], uint64(n)) // end of file
|
||||
binary.LittleEndian.PutUint64(f[48:], math.MaxUint64) // driver undefined
|
||||
binary.LittleEndian.PutUint64(f[64:], 96) // root object header
|
||||
return f
|
||||
}
|
||||
|
||||
// h5Msg is one object header message: its type and its body bytes.
|
||||
type h5Msg struct {
|
||||
typ uint16
|
||||
body []byte
|
||||
}
|
||||
|
||||
// h5ObjectHeader writes a version 1 object header at off carrying msgs
|
||||
// and returns the offset just past its message region.
|
||||
func h5ObjectHeader(f []byte, off int, msgs ...h5Msg) int {
|
||||
f[off] = 1
|
||||
binary.LittleEndian.PutUint16(f[off+2:], uint16(len(msgs)))
|
||||
binary.LittleEndian.PutUint32(f[off+4:], 1) // reference count
|
||||
region := off + 16
|
||||
size := 0
|
||||
for _, m := range msgs {
|
||||
size = alignUp(size+8+len(m.body), 8)
|
||||
}
|
||||
binary.LittleEndian.PutUint32(f[off+8:], uint32(size))
|
||||
for _, m := range msgs {
|
||||
binary.LittleEndian.PutUint16(f[region:], m.typ)
|
||||
binary.LittleEndian.PutUint16(f[region+2:], uint16(len(m.body)))
|
||||
copy(f[region+8:], m.body)
|
||||
region += alignUp(8+len(m.body), 8)
|
||||
}
|
||||
return region
|
||||
}
|
||||
|
||||
// h5Dataspace renders a version 1 dataspace message with the given
|
||||
// extents, each written as an eight-byte length.
|
||||
func h5Dataspace(dims ...uint64) []byte {
|
||||
m := make([]byte, 8+8*len(dims))
|
||||
m[0] = 1 // version 1
|
||||
if len(dims) == 0 {
|
||||
return m
|
||||
}
|
||||
m[1] = byte(len(dims))
|
||||
for i, d := range dims {
|
||||
binary.LittleEndian.PutUint64(m[8+8*i:], d)
|
||||
}
|
||||
return m
|
||||
}
|
||||
|
||||
// h5FloatType renders a version 1 floating-point datatype message of the
|
||||
// given element width.
|
||||
func h5FloatType(size uint32) []byte {
|
||||
m := make([]byte, 20)
|
||||
m[0] = 0x11 // version 1, class 1 (floating-point)
|
||||
binary.LittleEndian.PutUint32(m[4:], size)
|
||||
return m
|
||||
}
|
||||
|
||||
// h5ChunkLayoutV4 renders a version 4 chunked layout message whose chunk
|
||||
// dimensions are eight bytes wide, the flag the version 4 message
|
||||
// carries, so an extent beyond 2^32 can be declared at all.
|
||||
func h5ChunkLayoutV4(btree uint64, dims ...uint64) []byte {
|
||||
m := make([]byte, 3+8+1+8*len(dims))
|
||||
m[0] = 4
|
||||
m[1] = 2 // chunked
|
||||
m[2] = byte(len(dims))
|
||||
binary.LittleEndian.PutUint64(m[3:], btree)
|
||||
m[11] = 1 // eight-byte chunk dimensions
|
||||
for i, d := range dims {
|
||||
binary.LittleEndian.PutUint64(m[12+8*i:], d)
|
||||
}
|
||||
return m
|
||||
}
|
||||
|
||||
// h5ChunkLayoutV3 renders a version 3 chunked layout message with
|
||||
// four-byte chunk dimensions.
|
||||
func h5ChunkLayoutV3(btree uint64, dims ...uint32) []byte {
|
||||
m := make([]byte, 3+8+4*len(dims))
|
||||
m[0] = 3
|
||||
m[1] = 2
|
||||
m[2] = byte(len(dims))
|
||||
binary.LittleEndian.PutUint64(m[3:], btree)
|
||||
for i, d := range dims {
|
||||
binary.LittleEndian.PutUint32(m[11+4*i:], d)
|
||||
}
|
||||
return m
|
||||
}
|
||||
|
||||
// h5ContiguousLayout renders a version 3 contiguous layout message.
|
||||
func h5ContiguousLayout(addr, size uint64) []byte {
|
||||
m := make([]byte, 2+8+8)
|
||||
m[0] = 3
|
||||
m[1] = 1 // contiguous
|
||||
binary.LittleEndian.PutUint64(m[2:], addr)
|
||||
binary.LittleEndian.PutUint64(m[10:], size)
|
||||
return m
|
||||
}
|
||||
|
||||
// h5ChunkTree writes a one-entry leaf chunk B-tree at off for a dataset
|
||||
// of the given rank: one chunk of size stored bytes at chunkAt, with no
|
||||
// filter mask, at the given chunk offsets. The key carries one more
|
||||
// offset than the dataset has axes, the element size, which is why the
|
||||
// entry is 8+(8*(rank+1))+offSize bytes long.
|
||||
func h5ChunkTree(f []byte, off, rank int, size uint32, offsets []uint64, chunkAt uint64) {
|
||||
copy(f[off:], hdf5Tree)
|
||||
f[off+4] = 1 // chunk tree
|
||||
f[off+5] = 0 // leaf level
|
||||
binary.LittleEndian.PutUint16(f[off+6:], 1)
|
||||
p := off + 24
|
||||
binary.LittleEndian.PutUint32(f[p:], size)
|
||||
binary.LittleEndian.PutUint32(f[p+4:], 0) // filter mask
|
||||
for i := range rank {
|
||||
binary.LittleEndian.PutUint64(f[p+8+8*i:], offsets[i])
|
||||
}
|
||||
binary.LittleEndian.PutUint64(f[p+8+8*rank:], 8) // the element-size key
|
||||
binary.LittleEndian.PutUint64(f[p+8+8*(rank+1):], chunkAt)
|
||||
}
|
||||
|
||||
// TestLoadHDF5ChunkExtentWrap pins the chunk-dimension bound. The chunk
|
||||
// shape 2^32 by 2^29 float64 elements is 2^64 bytes, which wraps to
|
||||
// zero: the "chunk above the reader budget" guard and the "chunk holds a
|
||||
// full chunk" guard both saw 0, so the chunk strides were walked out of
|
||||
// the stored chunk (a slice out of range on a 2^61-element declaration,
|
||||
// and a silent read of the file's other bytes on a milder one).
|
||||
func TestLoadHDF5ChunkExtentWrap(t *testing.T) {
|
||||
cases := []struct {
|
||||
name string
|
||||
dims []uint64 // the chunk's shape plus the element-size slot
|
||||
}{
|
||||
{"2^32 by 2^29 elements", []uint64{1 << 32, 1 << 29, 8}},
|
||||
{"2^61 elements", []uint64{1 << 61, 1, 8}},
|
||||
}
|
||||
for _, tc := range cases {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
const btree, chunkAt, n = 320, 448, 512
|
||||
f := h5HostileFile(n)
|
||||
end := h5ObjectHeader(f, 96,
|
||||
h5Msg{hdf5MsgDataspace, h5Dataspace(2, 1)},
|
||||
h5Msg{hdf5MsgDatatype, h5FloatType(8)},
|
||||
h5Msg{hdf5MsgDataLayout, h5ChunkLayoutV4(btree, tc.dims...)},
|
||||
)
|
||||
if end > btree {
|
||||
t.Fatalf("the test object header runs to %d, past the chunk B-tree at %d", end, btree)
|
||||
}
|
||||
h5ChunkTree(f, btree, 2, 16, []uint64{0, 0}, chunkAt)
|
||||
path := writeHostile(t, "chunkextent.h5", f)
|
||||
sets, err := LoadHDF5(path)
|
||||
if err == nil {
|
||||
t.Fatalf("LoadHDF5 accepted a chunk whose extent wraps: %d datasets", len(sets))
|
||||
}
|
||||
if !strings.Contains(err.Error(), "budget") {
|
||||
t.Fatalf("error = %v, want the reader-budget refusal", err)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// TestLoadHDF5DatasetExtentWrap pins the dataspace bound on both storage
|
||||
// classes: a shape of 2^32 by 2^32 float64 elements wraps the byte count
|
||||
// to zero, which passed the cap, so the output buffer was allocated
|
||||
// empty while the copy needed a whole element (the chunked class
|
||||
// panicked), and the contiguous class returned a dataset claiming 2^64
|
||||
// elements with an empty payload and no error.
|
||||
func TestLoadHDF5DatasetExtentWrap(t *testing.T) {
|
||||
t.Run("chunked", func(t *testing.T) {
|
||||
const btree, chunkAt, n = 320, 448, 512
|
||||
f := h5HostileFile(n)
|
||||
end := h5ObjectHeader(f, 96,
|
||||
h5Msg{hdf5MsgDataspace, h5Dataspace(1<<32, 1<<32)},
|
||||
h5Msg{hdf5MsgDatatype, h5FloatType(8)},
|
||||
h5Msg{hdf5MsgDataLayout, h5ChunkLayoutV3(btree, 1, 1, 8)},
|
||||
)
|
||||
if end > btree {
|
||||
t.Fatalf("the test object header runs to %d, past the chunk B-tree at %d", end, btree)
|
||||
}
|
||||
h5ChunkTree(f, btree, 2, 8, []uint64{0, 0}, chunkAt)
|
||||
sets, err := LoadHDF5(writeHostile(t, "shapedwrap.h5", f))
|
||||
if err == nil {
|
||||
t.Fatalf("LoadHDF5 accepted a wrapping dataspace: %d datasets", len(sets))
|
||||
}
|
||||
if !strings.Contains(err.Error(), "budget") {
|
||||
t.Fatalf("error = %v, want the reader-budget refusal", err)
|
||||
}
|
||||
})
|
||||
t.Run("contiguous", func(t *testing.T) {
|
||||
const dataAt, n = 448, 512
|
||||
f := h5HostileFile(n)
|
||||
h5ObjectHeader(f, 96,
|
||||
h5Msg{hdf5MsgDataspace, h5Dataspace(1<<32, 1<<32)},
|
||||
h5Msg{hdf5MsgDatatype, h5FloatType(8)},
|
||||
h5Msg{hdf5MsgDataLayout, h5ContiguousLayout(dataAt, 8)},
|
||||
)
|
||||
sets, err := LoadHDF5(writeHostile(t, "contigwrap.h5", f))
|
||||
if err == nil {
|
||||
shape := []int(nil)
|
||||
if len(sets) > 0 {
|
||||
shape = sets[0].Shape
|
||||
}
|
||||
t.Fatalf("LoadHDF5 returned %v for a wrapping dataspace, want an error", shape)
|
||||
}
|
||||
if !strings.Contains(err.Error(), "budget") {
|
||||
t.Fatalf("error = %v, want the reader-budget refusal", err)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
// TestLoadHDF5AttributeExtentWrap pins the attribute dataspace bound: a
|
||||
// 3 by 2^62 element dataspace wraps its element count to a negative
|
||||
// number in the multiply-first check, which let it past the bounds test
|
||||
// and into an allocation with a negative capacity. The attribute must be
|
||||
// refused (dropped), not accepted with a wrapped value.
|
||||
func TestLoadHDF5AttributeExtentWrap(t *testing.T) {
|
||||
const n = 512
|
||||
f := h5HostileFile(n)
|
||||
// The attribute message: version, flags, then the three size fields
|
||||
// with their fields behind them. A 2^62 extent is far past the
|
||||
// message, so no legal value can follow it.
|
||||
attr := make([]byte, 88)
|
||||
attr[0] = 1
|
||||
binary.LittleEndian.PutUint16(attr[2:], 2) // name size
|
||||
binary.LittleEndian.PutUint16(attr[4:], 20) // datatype size
|
||||
binary.LittleEndian.PutUint16(attr[6:], 24) // dataspace size
|
||||
copy(attr[8:], "n\x00")
|
||||
copy(attr[16:], h5FloatType(8))
|
||||
copy(attr[40:], h5Dataspace(3, 1<<62))
|
||||
// A hard link to a one-element dataset, so the attribute's fate is
|
||||
// observable: an accepted attribute lands in the dataset's map.
|
||||
link := func(addr uint64) []byte {
|
||||
return binary.LittleEndian.AppendUint64([]byte{1, 0, 1, 'd'}, addr)
|
||||
}
|
||||
datasetAt := 288
|
||||
end := h5ObjectHeader(f, 96,
|
||||
h5Msg{hdf5MsgLink, link(uint64(datasetAt))},
|
||||
h5Msg{hdf5MsgAttribute, attr},
|
||||
)
|
||||
if end > datasetAt {
|
||||
t.Fatalf("the test object header runs to %d, past the dataset at %d", end, datasetAt)
|
||||
}
|
||||
h5ObjectHeader(f, datasetAt,
|
||||
h5Msg{hdf5MsgDataspace, h5Dataspace(1)},
|
||||
h5Msg{hdf5MsgDatatype, h5FloatType(8)},
|
||||
h5Msg{hdf5MsgDataLayout, h5ContiguousLayout(448, 8)},
|
||||
)
|
||||
sets, err := LoadHDF5(writeHostile(t, "attrextent.h5", f))
|
||||
if err != nil {
|
||||
t.Fatalf("LoadHDF5 refused the file: %v", err)
|
||||
}
|
||||
if len(sets) != 1 || sets[0].Path != "/d" {
|
||||
t.Fatalf("datasets = %d, want the linked /d", len(sets))
|
||||
}
|
||||
if v, ok := sets[0].Attrs["n"]; ok {
|
||||
t.Fatalf("the attribute with a wrapping dataspace was accepted as %q", v)
|
||||
}
|
||||
}
|
||||
|
||||
// TestLoadHDF5MessageCountBounded pins the object header preallocation:
|
||||
// the slice used to be sized from the header's declared message count
|
||||
// (65535 messages of 32 bytes is 2 MiB) before the region that would
|
||||
// have to hold them was consulted, so 16 bytes of file ordered a
|
||||
// two-megabyte allocation.
|
||||
func TestLoadHDF5MessageCountBounded(t *testing.T) {
|
||||
f := h5HostileFile(160)
|
||||
f[96] = 1
|
||||
binary.LittleEndian.PutUint16(f[98:], 65535) // declared messages
|
||||
binary.LittleEndian.PutUint32(f[100:], 1) // reference count
|
||||
binary.LittleEndian.PutUint32(f[104:], 16) // the region that holds them
|
||||
path := writeHostile(t, "msgcount.h5", f)
|
||||
runtime.GC()
|
||||
var before, after runtime.MemStats
|
||||
runtime.ReadMemStats(&before)
|
||||
_, err := LoadHDF5(path)
|
||||
runtime.ReadMemStats(&after)
|
||||
if err == nil {
|
||||
t.Fatal("LoadHDF5 accepted a header whose messages run past the region")
|
||||
}
|
||||
if used := after.TotalAlloc - before.TotalAlloc; used > 512<<10 {
|
||||
t.Fatalf("loading a 160-byte file allocated %d bytes: the declared message count still sizes the preallocation", used)
|
||||
}
|
||||
}
|
||||
|
||||
// TestLoadNetCDFRecordSizeWrap pins the record-size pre-pass. It used to
|
||||
// multiply every record variable's non-record dimensions before any
|
||||
// variable was validated, so two hostile declarations inflated the
|
||||
// record size for an innocuous variable to 2^54 bytes: one element per
|
||||
// record then made (numrecs-1)*recordSize exactly 2^64, which wrapped to
|
||||
// zero, the span check passed and the reader walked 2^54 bytes into a
|
||||
// 4 KiB file.
|
||||
func TestLoadNetCDFRecordSizeWrap(t *testing.T) {
|
||||
const numrecs = 1<<10 + 1
|
||||
// The arithmetic the reader performed: two slabs of 2^25 by 2^25
|
||||
// doubles are 2^53 bytes each and the innocuous variable's slab is
|
||||
// four, so the record size is 2^54+4 and the last record's slab
|
||||
// offset wraps to 2^64+4096, i.e. 4096. A file larger than that
|
||||
// passes the span check and the following records then index far
|
||||
// past the file.
|
||||
recordSize := 2*int64(1<<53) + 4
|
||||
if span := (numrecs-1)*recordSize + 1; span != 1<<12+1 {
|
||||
t.Fatalf("construction is wrong: the wrapped span is %d, want %d", span, 1<<12+1)
|
||||
}
|
||||
|
||||
var b []byte
|
||||
b = append(b, 'C', 'D', 'F', 1)
|
||||
b = binary.BigEndian.AppendUint32(b, numrecs)
|
||||
// dim_list: the record dimension leads, then one one-element
|
||||
// dimension and the two hostile ones.
|
||||
b = binary.BigEndian.AppendUint32(b, ncTagDimension)
|
||||
b = binary.BigEndian.AppendUint32(b, 4)
|
||||
b = hostileNCName(b, "rec")
|
||||
b = binary.BigEndian.AppendUint32(b, 0)
|
||||
b = hostileNCName(b, "one")
|
||||
b = binary.BigEndian.AppendUint32(b, 1)
|
||||
b = hostileNCName(b, "big1")
|
||||
b = binary.BigEndian.AppendUint32(b, 1<<25)
|
||||
b = hostileNCName(b, "big2")
|
||||
b = binary.BigEndian.AppendUint32(b, 1<<25)
|
||||
// gatt_list: absent.
|
||||
b = binary.BigEndian.AppendUint32(b, 0)
|
||||
b = binary.BigEndian.AppendUint32(b, 0)
|
||||
// var_list: the innocuous record variable, then the two that inflate
|
||||
// the record size.
|
||||
b = binary.BigEndian.AppendUint32(b, ncTagVariable)
|
||||
b = binary.BigEndian.AppendUint32(b, 3)
|
||||
b = hostileNCName(b, "a")
|
||||
b = binary.BigEndian.AppendUint32(b, 2) // rank: rec, one
|
||||
b = binary.BigEndian.AppendUint32(b, 0)
|
||||
b = binary.BigEndian.AppendUint32(b, 1)
|
||||
b = binary.BigEndian.AppendUint32(b, 0) // variable attributes: absent
|
||||
b = binary.BigEndian.AppendUint32(b, 0)
|
||||
b = binary.BigEndian.AppendUint32(b, ncTypeByte)
|
||||
b = binary.BigEndian.AppendUint32(b, 4) // vsize
|
||||
b = binary.BigEndian.AppendUint32(b, 0) // begin
|
||||
for i := range 2 {
|
||||
b = hostileNCName(b, fmt.Sprintf("b%d", i))
|
||||
b = binary.BigEndian.AppendUint32(b, 3) // rank: rec, big1, big2
|
||||
b = binary.BigEndian.AppendUint32(b, 0)
|
||||
b = binary.BigEndian.AppendUint32(b, 2)
|
||||
b = binary.BigEndian.AppendUint32(b, 3)
|
||||
b = binary.BigEndian.AppendUint32(b, 0) // variable attributes: absent
|
||||
b = binary.BigEndian.AppendUint32(b, 0)
|
||||
b = binary.BigEndian.AppendUint32(b, ncTypeDouble)
|
||||
b = binary.BigEndian.AppendUint32(b, 0) // vsize
|
||||
b = binary.BigEndian.AppendUint32(b, 0) // begin
|
||||
}
|
||||
// A file long enough for the innocuous variable's own records and for
|
||||
// the wrapped span the buggy check computes.
|
||||
for len(b) < 8192 {
|
||||
b = append(b, 0)
|
||||
}
|
||||
path := writeHostile(t, "recsize-wrap.nc", b)
|
||||
_, _, _, err := LoadNetCDF(path)
|
||||
if err == nil {
|
||||
t.Fatal("LoadNetCDF accepted a record variable whose dimensions cannot fit the file")
|
||||
}
|
||||
if !strings.Contains(err.Error(), "more elements than the file holds") {
|
||||
t.Fatalf("error = %v, want the per-variable element bound", err)
|
||||
}
|
||||
}
|
||||
|
||||
// TestSaveFITSTableASCIIFloat32 pins the ASCII float encoder for a
|
||||
// float32 column. The encoder read the float64 payload first and only
|
||||
// then substituted the float32 one, but RawFloats is nil for every
|
||||
// float32 array, so the first row of an ordinary float32 column indexed
|
||||
// a nil slice.
|
||||
func TestSaveFITSTableASCIIFloat32(t *testing.T) {
|
||||
vals := []float32{1.5, -2.25, 3e8}
|
||||
f32 := core.New(core.Float32, len(vals))
|
||||
copy(f32.RawFloat32s(), vals)
|
||||
cols := []FITSTableColumn{{Name: "F", Form: "E12.4", Data: f32}}
|
||||
path := filepath.Join(t.TempDir(), "f32ascii.fits")
|
||||
if err := SaveFITSTable(path, true, cols, nil); err != nil {
|
||||
t.Fatalf("SaveFITSTable: %v", err)
|
||||
}
|
||||
table, err := LoadFITSTable(path)
|
||||
if err != nil {
|
||||
t.Fatalf("LoadFITSTable: %v", err)
|
||||
}
|
||||
for i, want := range vals {
|
||||
if got := table.Columns[0].FloatAt(i); got != float64(want) {
|
||||
t.Fatalf("F[%d] = %g, want %g", i, got, float64(want))
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// TestLoadFITSTableFortranNumbers pins the number forms a Fortran writer
|
||||
// puts in an ASCII table: strconv.ParseFloat accepts neither the Dw.d
|
||||
// exponent ("1.5D+03") nor the exponent-less form ("1.5+03"), and one
|
||||
// unparseable cell refused the whole table.
|
||||
func TestLoadFITSTableFortranNumbers(t *testing.T) {
|
||||
hdr := cardBlock(
|
||||
card("XTENSION= 'TABLE '"),
|
||||
card("BITPIX = 8"),
|
||||
card("NAXIS = 2"),
|
||||
card("NAXIS1 = 20"),
|
||||
card("NAXIS2 = 4"),
|
||||
card("TFIELDS = 1"),
|
||||
card("TBCOL1 = 1"),
|
||||
card("TFORM1 = 'D20.12'"),
|
||||
card("END"),
|
||||
)
|
||||
rows := []string{"1.5D+03", "1.5+03", "-2.5d-2", "1.5E+03"}
|
||||
payload := make([]byte, 0, 20*len(rows))
|
||||
for _, r := range rows {
|
||||
payload = append(payload, fmt.Sprintf("%-20s", r)...)
|
||||
}
|
||||
path := writeHostile(t, "fortran.fits", append(hdr, payload...))
|
||||
table, err := LoadFITSTable(path)
|
||||
if err != nil {
|
||||
t.Fatalf("LoadFITSTable refused a Fortran table: %v", err)
|
||||
}
|
||||
want := []float64{1500, 1500, -0.025, 1500}
|
||||
for i, w := range want {
|
||||
if got := table.Columns[0].FloatAt(i); got != w {
|
||||
t.Fatalf("row %d (%q) = %g, want %g", i, rows[i], got, w)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// TestLoadFITSTableScaleOverflow pins the integral-scaling bound. The
|
||||
// guard used 9.3e18, which is above MaxInt64 (9.223372036854775807e18),
|
||||
// so a scaled value in [2^63, 9.3e18) reached an out-of-range
|
||||
// float-to-integer conversion: the column kept the int dtype and the
|
||||
// value came back as MinInt64 instead of the documented promotion to
|
||||
// float64.
|
||||
func TestLoadFITSTableScaleOverflow(t *testing.T) {
|
||||
const raw = int64(1) << 62
|
||||
cols := []FITSTableColumn{{Name: "V", Form: "K", Data: mustInts(t, []int64{raw}, 1)}}
|
||||
path := filepath.Join(t.TempDir(), "scaled.fits")
|
||||
if err := SaveFITSTable(path, false, cols, map[string]string{"TSCAL1": "2", "TZERO1": "0"}); err != nil {
|
||||
t.Fatalf("SaveFITSTable: %v", err)
|
||||
}
|
||||
table, err := LoadFITSTable(path)
|
||||
if err != nil {
|
||||
t.Fatalf("LoadFITSTable: %v", err)
|
||||
}
|
||||
col := table.Columns[0]
|
||||
want := float64(raw) * 2 // 2^63, exactly representable in float64
|
||||
if col.Dtype() != core.Float {
|
||||
t.Fatalf("dtype = %s, want float64: a scaling past MaxInt64 must promote the column", col.Dtype())
|
||||
}
|
||||
if got := col.FloatAt(0); got != want {
|
||||
t.Fatalf("scaled value = %g, want %g", got, want)
|
||||
}
|
||||
}
|
||||
|
||||
// TestMapCountsWithoutWrap pins the mmap count arithmetic: the byte
|
||||
// length used to be int64(n)*elementSize, which wraps for a hostile
|
||||
// count, so the wrapped length passed the file-size checks and the typed
|
||||
// view was then built with the original count, which unsafe.Slice
|
||||
// rejects with a panic instead of an error.
|
||||
func TestMapCountsWithoutWrap(t *testing.T) {
|
||||
path := filepath.Join(t.TempDir(), "small.bin")
|
||||
if err := SaveNativeFloats(path, []float64{1, 2, 3, 4}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
// 2^62+3 float64 values are 2^65+24 bytes, which wraps to 24: the
|
||||
// file holds 32. The others already failed the non-positive-length
|
||||
// check by luck and must keep failing it.
|
||||
for _, n := range []int{1<<62 + 3, 1 << 61, 1 << 60, math.MaxInt} {
|
||||
if _, _, err := MapFloats(path, 0, n); err == nil {
|
||||
t.Errorf("MapFloats(n = %d) accepted a count the file cannot hold", n)
|
||||
}
|
||||
if _, _, err := MapFloat32s(path, 0, n); err == nil {
|
||||
t.Errorf("MapFloat32s(n = %d) accepted a count the file cannot hold", n)
|
||||
}
|
||||
if _, _, err := MapInts(path, 0, n); err == nil {
|
||||
t.Errorf("MapInts(n = %d) accepted a count the file cannot hold", n)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// TestSaveCSVZeroColumns pins the shape contract: an array with rows and
|
||||
// no columns wrote one empty record per row, and encoding/csv reads
|
||||
// blank lines as no records at all, so (3,0) came back as (0,0). It is
|
||||
// refused instead, because no CSV text can carry the column count.
|
||||
func TestSaveCSVZeroColumns(t *testing.T) {
|
||||
a, err := core.FromFloats(nil, 3, 0)
|
||||
if err != nil {
|
||||
t.Fatalf("FromFloats(nil, 3, 0): %v", err)
|
||||
}
|
||||
var sb strings.Builder
|
||||
if err := SaveCSVWriter(&sb, a); err == nil {
|
||||
t.Fatalf("SaveCSVWriter accepted the shape %v and wrote %q, which reads back as a different shape",
|
||||
a.Shape(), sb.String())
|
||||
}
|
||||
// A 2-D array with columns still round-trips.
|
||||
ok := mustFloats(t, []float64{1, 2, 3, 4, 5, 6}, 2, 3)
|
||||
var back strings.Builder
|
||||
if err := SaveCSVWriter(&back, ok); err != nil {
|
||||
t.Fatalf("SaveCSVWriter: %v", err)
|
||||
}
|
||||
got, err := LoadCSVReader(strings.NewReader(back.String()), false)
|
||||
if err != nil {
|
||||
t.Fatalf("LoadCSVReader: %v", err)
|
||||
}
|
||||
if got.Shape()[0] != 2 || got.Shape()[1] != 3 {
|
||||
t.Fatalf("round trip gave %v, want [2 3]", got.Shape())
|
||||
}
|
||||
}
|
||||
|
||||
// TestIOSourceTypography pins the plain-ASCII rule for the package's own
|
||||
// source: an em dash, an en dash, a Unicode minus, a middle dot or a
|
||||
// multiplication sign in a comment or a message is a typographic symbol
|
||||
// where the repository's style asks for ASCII.
|
||||
func TestIOSourceTypography(t *testing.T) {
|
||||
banned := []struct {
|
||||
r rune
|
||||
name string
|
||||
}{
|
||||
{'\u2014', "em dash"},
|
||||
{'\u2013', "en dash"},
|
||||
{'\u2212', "Unicode minus"},
|
||||
{'\u00b7', "middle dot"},
|
||||
{'\u00d7', "multiplication sign"},
|
||||
}
|
||||
entries, err := os.ReadDir(".")
|
||||
if err != nil {
|
||||
t.Fatalf("ReadDir: %v", err)
|
||||
}
|
||||
checked := 0
|
||||
for _, e := range entries {
|
||||
name := e.Name()
|
||||
if !strings.HasSuffix(name, ".go") || strings.HasSuffix(name, "_test.go") {
|
||||
continue
|
||||
}
|
||||
data, err := os.ReadFile(name)
|
||||
if err != nil {
|
||||
t.Fatalf("ReadFile(%s): %v", name, err)
|
||||
}
|
||||
checked++
|
||||
for _, b := range banned {
|
||||
if strings.ContainsRune(string(data), b.r) {
|
||||
t.Errorf("%s carries a %s (U+%04X)", name, b.name, b.r)
|
||||
}
|
||||
}
|
||||
}
|
||||
if checked == 0 {
|
||||
t.Fatal("no source files were checked")
|
||||
}
|
||||
}
|
||||
+561
@@ -0,0 +1,561 @@
|
||||
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
|
||||
// SPDX-License-Identifier: MIT
|
||||
|
||||
package io
|
||||
|
||||
import (
|
||||
"encoding/binary"
|
||||
"fmt"
|
||||
"maps"
|
||||
"math"
|
||||
"os"
|
||||
"slices"
|
||||
"strconv"
|
||||
"strings"
|
||||
|
||||
"sourcedock.dev/petrbalvin/tensor/internal/base"
|
||||
"sourcedock.dev/petrbalvin/tensor/internal/core"
|
||||
)
|
||||
|
||||
// FITS image I/O. The Flexible Image Transport System is astronomy's
|
||||
// archival format: a self-describing header of 80-character ASCII
|
||||
// cards in 2880-byte blocks, followed by big-endian binary data
|
||||
// padded to the same block size. Every observation archived by a
|
||||
// telescope in the last four decades reads back with the same parser,
|
||||
// which is the property that makes a format worth speaking.
|
||||
//
|
||||
// This implementation covers the primary HDU image: BITPIX -64
|
||||
// (float64) and -32 (float32), any rank with positive axes. FITS
|
||||
// orders axes Fortran-style with NAXIS1 varying fastest, the opposite
|
||||
// of Go's row-major convention, so NAXISj is declared from the shape
|
||||
// in reverse and the flat payload needs no permutation. Extensions
|
||||
// (XTENSION), tables and the integer bit depths are refused with an
|
||||
// error rather than half-read.
|
||||
|
||||
// fitsCardsPerBlock is the number of 80-byte cards in one 2880-byte
|
||||
// FITS block.
|
||||
const fitsCardsPerBlock = 2880 / 80
|
||||
|
||||
// SaveFITS writes a float64 or float32 array as a FITS primary image
|
||||
// with the given header entries. Keywords are uppercased and must be
|
||||
// 1 to 8 characters from A-Z, 0-9, '-' and '_' (the format's reserved
|
||||
// SIMPLE, BITPIX, NAXIS, NAXISn, EXTEND and END are refused); values
|
||||
// are written as FITS strings, at most 68 characters after the
|
||||
// format's quote escaping.
|
||||
func SaveFITS(path string, a *core.Array, headers map[string]string) error {
|
||||
var bitpix int
|
||||
switch a.Dtype() {
|
||||
case core.Float:
|
||||
bitpix = -64
|
||||
case core.Float32:
|
||||
bitpix = -32
|
||||
default:
|
||||
return base.Errf("SaveFITS: supports float64 and float32 arrays, got dtype %s", a.Dtype())
|
||||
}
|
||||
if a.NDim() == 0 {
|
||||
return base.Errf("SaveFITS: the image needs at least one axis")
|
||||
}
|
||||
for i, d := range a.Shape() {
|
||||
if d <= 0 {
|
||||
return base.Errf("SaveFITS: axis %d has extent %d, every axis must be positive", i+1, d)
|
||||
}
|
||||
}
|
||||
|
||||
cards := []string{
|
||||
fitsBoolCard("SIMPLE", true),
|
||||
fitsIntCard("BITPIX", bitpix),
|
||||
fitsIntCard("NAXIS", a.NDim()),
|
||||
}
|
||||
// NAXIS1 is the fastest-varying axis, the last one in Go's
|
||||
// row-major order.
|
||||
for j := range a.NDim() {
|
||||
cards = append(cards, fitsIntCard("NAXIS"+strconv.Itoa(j+1), a.Shape()[a.NDim()-1-j]))
|
||||
}
|
||||
cards = append(cards, fitsBoolCard("EXTEND", true))
|
||||
userCards, err := fitsUserCards(headers)
|
||||
if err != nil {
|
||||
return base.Errf("SaveFITS: %w", err)
|
||||
}
|
||||
cards = append(cards, userCards...)
|
||||
cards = append(cards, fitsEndCard())
|
||||
|
||||
elem := a.Len()
|
||||
width := 8
|
||||
if a.Dtype() == core.Float32 {
|
||||
width = 4
|
||||
}
|
||||
// The final size is known up front: both the header and the payload
|
||||
// pad to whole blocks, so one allocation serves the whole file.
|
||||
out := make([]byte, 0, fitsBlockSize(len(cards))+fitsBlockSize(elem*width))
|
||||
out = fitsAppendCards(out, cards)
|
||||
if a.Dtype() == core.Float {
|
||||
raw := a.RawFloats()
|
||||
for i := range elem {
|
||||
out = binary.BigEndian.AppendUint64(out, math.Float64bits(raw[i]))
|
||||
}
|
||||
} else {
|
||||
raw := a.RawFloat32s()
|
||||
for i := range elem {
|
||||
out = binary.BigEndian.AppendUint32(out, math.Float32bits(raw[i]))
|
||||
}
|
||||
}
|
||||
// Zero bytes pad the data to the block boundary.
|
||||
out = fitsAppendZeroPad(out)
|
||||
return os.WriteFile(path, out, 0o644)
|
||||
}
|
||||
|
||||
// LoadFITS reads a FITS primary image into a float64 (BITPIX -64) or
|
||||
// float32 (BITPIX -32) array, returning every non-structural header
|
||||
// entry alongside it. String values are unquoted and unescaped,
|
||||
// logical values come back as "T" or "F", numbers as their literal
|
||||
// text; COMMENT, HISTORY and blank cards carry no value and are
|
||||
// skipped.
|
||||
func LoadFITS(path string) (*core.Array, map[string]string, error) {
|
||||
data, err := os.ReadFile(path)
|
||||
if err != nil {
|
||||
return nil, nil, base.Errf("LoadFITS: %w", err)
|
||||
}
|
||||
return parseFITS(data)
|
||||
}
|
||||
|
||||
// parseFITS decodes a FITS primary image from raw bytes.
|
||||
func parseFITS(data []byte) (*core.Array, map[string]string, error) {
|
||||
if len(data) < 80 {
|
||||
return nil, nil, base.Errf("LoadFITS: file is shorter than one header card")
|
||||
}
|
||||
if string(data[:9]) == "XTENSION " {
|
||||
return nil, nil, base.Errf("LoadFITS: extensions are not supported, only the primary image")
|
||||
}
|
||||
if strings.TrimRight(string(data[:8]), " ") != "SIMPLE" {
|
||||
return nil, nil, base.Errf("LoadFITS: the first card must be SIMPLE")
|
||||
}
|
||||
|
||||
cards, dataAt, err := scanFITSCards(data, 0)
|
||||
if err != nil {
|
||||
return nil, nil, base.Errf("LoadFITS: %w", err)
|
||||
}
|
||||
|
||||
var (
|
||||
headers = map[string]string{}
|
||||
bitpix int
|
||||
naxis = -1
|
||||
axisVals = map[int]int{}
|
||||
dims []int
|
||||
)
|
||||
for _, c := range cards {
|
||||
switch {
|
||||
case c.key == "SIMPLE":
|
||||
if c.value != "T" {
|
||||
return nil, nil, base.Errf("LoadFITS: SIMPLE = F marks a non-conformant file")
|
||||
}
|
||||
case c.key == "EXTEND":
|
||||
// Structural; EXTEND still reports itself to the caller.
|
||||
headers[c.key] = c.value
|
||||
case c.key == "BITPIX":
|
||||
v, cerr := strconv.Atoi(c.value)
|
||||
if cerr != nil {
|
||||
return nil, nil, base.Errf("LoadFITS: %s = %q is not an integer", c.key, c.value)
|
||||
}
|
||||
bitpix = v
|
||||
case c.key == "NAXIS":
|
||||
v, cerr := strconv.Atoi(c.value)
|
||||
if cerr != nil {
|
||||
return nil, nil, base.Errf("LoadFITS: %s = %q is not an integer", c.key, c.value)
|
||||
}
|
||||
naxis = v
|
||||
default:
|
||||
// NAXISn is an axis only when a number follows the prefix,
|
||||
// the rule fitsCheckKeyword applies on write kept symmetric;
|
||||
// NAXISREF and every other spelling is a user keyword. The
|
||||
// value is stored under its axis number, so the card order
|
||||
// cannot re-bind the axes.
|
||||
if fitsAxisKeyword(c.key) {
|
||||
v, cerr := strconv.Atoi(c.value)
|
||||
if cerr != nil {
|
||||
return nil, nil, base.Errf("LoadFITS: %s = %q is not an integer", c.key, c.value)
|
||||
}
|
||||
j, _ := strconv.Atoi(c.key[len("NAXIS"):])
|
||||
if j < 1 {
|
||||
return nil, nil, base.Errf("LoadFITS: %s is not an axis keyword", c.key)
|
||||
}
|
||||
if _, dup := axisVals[j]; dup {
|
||||
return nil, nil, base.Errf("LoadFITS: %s repeats", c.key)
|
||||
}
|
||||
axisVals[j] = v
|
||||
continue
|
||||
}
|
||||
headers[c.key] = c.value
|
||||
}
|
||||
}
|
||||
// A zero-axis primary HDU (the standard container for extension
|
||||
// files) answers an empty image.
|
||||
if naxis == 0 {
|
||||
return core.New(core.Float, 0), headers, nil
|
||||
}
|
||||
if bitpix != -64 && bitpix != -32 {
|
||||
return nil, nil, base.Errf("LoadFITS: BITPIX %d is not supported (want -64 or -32)", bitpix)
|
||||
}
|
||||
// NAXISn values are bound by their axis number, not by card order:
|
||||
// a header writing NAXIS2 before NAXIS1, or omitting an axis, is
|
||||
// malformed and must be refused rather than silently re-bound.
|
||||
if naxis < 0 {
|
||||
return nil, nil, base.Errf("LoadFITS: NAXIS = %d is negative or the card is missing", naxis)
|
||||
}
|
||||
// The card count gates the allocation: a hostile NAXIS far beyond
|
||||
// the NAXISn cards the file actually carries is refused here, not
|
||||
// turned into a slice of that length.
|
||||
if len(axisVals) != naxis {
|
||||
return nil, nil, base.Errf("LoadFITS: NAXIS = %d with %d NAXISn cards", naxis, len(axisVals))
|
||||
}
|
||||
dims = make([]int, naxis)
|
||||
for j := 1; j <= naxis; j++ {
|
||||
d, ok := axisVals[j]
|
||||
if !ok {
|
||||
return nil, nil, base.Errf("LoadFITS: NAXIS%d is missing under NAXIS = %d", j, naxis)
|
||||
}
|
||||
dims[j-1] = d
|
||||
}
|
||||
width := -bitpix / 8
|
||||
if dataAt > len(data) {
|
||||
return nil, nil, base.Errf("LoadFITS: data is truncated (%d bytes present)", len(data))
|
||||
}
|
||||
// Every axis extent and running product is bounded by the bytes
|
||||
// the file actually holds, so a hostile header cannot overflow the
|
||||
// int product before the truncation check rejects it.
|
||||
avail := (len(data) - dataAt) / width
|
||||
shape := make([]int, naxis)
|
||||
total := 1
|
||||
for i, d := range dims {
|
||||
if d <= 0 {
|
||||
return nil, nil, base.Errf("LoadFITS: NAXIS%d = %d, every axis must be positive", i+1, d)
|
||||
}
|
||||
if d > avail || total > avail/d {
|
||||
return nil, nil, base.Errf("LoadFITS: data is truncated (%d bytes present, more needed)",
|
||||
len(data)-dataAt)
|
||||
}
|
||||
shape[naxis-1-i] = d
|
||||
total *= d
|
||||
}
|
||||
|
||||
payload := data[dataAt : dataAt+total*width]
|
||||
var out *core.Array
|
||||
if bitpix == -64 {
|
||||
out = core.New(core.Float, shape...)
|
||||
raw := out.RawFloats()
|
||||
for i := range total {
|
||||
raw[i] = math.Float64frombits(binary.BigEndian.Uint64(payload[i*8:]))
|
||||
}
|
||||
} else {
|
||||
out = core.New(core.Float32, shape...)
|
||||
raw := out.RawFloat32s()
|
||||
for i := range total {
|
||||
raw[i] = math.Float32frombits(binary.BigEndian.Uint32(payload[i*4:]))
|
||||
}
|
||||
}
|
||||
// BSCALE/BZERO scaling: physical = raw*scale + zero (FITS 4.1).
|
||||
// A silent skip would hand back storage values as physical ones, so
|
||||
// the affine map is applied whenever the keywords deviate from the
|
||||
// identity; float32 values are computed in float64 and rounded once.
|
||||
scale, serr := fitsScaledHeader(headers, "BSCALE", 1)
|
||||
if serr != nil {
|
||||
return nil, nil, base.Errf("LoadFITS: %w", serr)
|
||||
}
|
||||
zero, zerr := fitsScaledHeader(headers, "BZERO", 0)
|
||||
if zerr != nil {
|
||||
return nil, nil, base.Errf("LoadFITS: %w", zerr)
|
||||
}
|
||||
if scale != 1 || zero != 0 {
|
||||
if out.Dtype() == core.Float {
|
||||
raw := out.RawFloats()
|
||||
for i := range raw {
|
||||
raw[i] = raw[i]*scale + zero
|
||||
}
|
||||
} else {
|
||||
raw := out.RawFloat32s()
|
||||
for i := range raw {
|
||||
raw[i] = float32(float64(raw[i])*scale + zero)
|
||||
}
|
||||
}
|
||||
}
|
||||
return out, headers, nil
|
||||
}
|
||||
|
||||
// fitsScaledHeader parses a floating-point header entry, falling back
|
||||
// to def when the keyword is absent. A present but malformed value is
|
||||
// an error, never silently the default.
|
||||
func fitsScaledHeader(headers map[string]string, key string, def float64) (float64, error) {
|
||||
v, ok := headers[key]
|
||||
if !ok || v == "" {
|
||||
return def, nil
|
||||
}
|
||||
f, err := strconv.ParseFloat(strings.TrimSpace(v), 64)
|
||||
if err != nil {
|
||||
return 0, base.Errf("%s = %q is not a number", key, v)
|
||||
}
|
||||
return f, nil
|
||||
}
|
||||
|
||||
// fitsValue extracts the value field of a card: everything after the
|
||||
// "= " indicator, minus any trailing comment, with FITS string
|
||||
// quoting resolved. The second return says whether the value was a
|
||||
// quoted string, which the long-string CONTINUE convention keys on.
|
||||
func fitsValue(card string) (string, bool, error) {
|
||||
field := card[10:]
|
||||
if strings.HasPrefix(strings.TrimLeft(field, " "), "'") {
|
||||
// A quoted string: '' inside escapes one quote.
|
||||
var b strings.Builder
|
||||
in := field[strings.Index(field, "'"):]
|
||||
i := 1
|
||||
for i < len(in) {
|
||||
if in[i] == '\'' {
|
||||
if i+1 < len(in) && in[i+1] == '\'' {
|
||||
b.WriteByte('\'')
|
||||
i += 2
|
||||
continue
|
||||
}
|
||||
return strings.TrimRight(b.String(), " "), true, nil
|
||||
}
|
||||
b.WriteByte(in[i])
|
||||
i++
|
||||
}
|
||||
return "", false, base.Errf("card %q has an unterminated string value", card[:min(20, len(card))])
|
||||
}
|
||||
// Free-format value: cut at the comment slash and trim.
|
||||
if slash := strings.IndexByte(field, '/'); slash >= 0 {
|
||||
field = field[:slash]
|
||||
}
|
||||
return strings.TrimSpace(field), false, nil
|
||||
}
|
||||
|
||||
// fitsContinueString extracts the quoted segment of a CONTINUE card.
|
||||
// The keyword occupies columns 1-8 and no "= " indicator follows, so
|
||||
// the string opens at the first quote anywhere in the card.
|
||||
func fitsContinueString(card string) (string, error) {
|
||||
q := strings.IndexByte(card, '\'')
|
||||
if q < 0 {
|
||||
return "", base.Errf("CONTINUE card %q has no quoted segment", card[:min(20, len(card))])
|
||||
}
|
||||
var b strings.Builder
|
||||
i := q + 1
|
||||
for i < len(card) {
|
||||
if card[i] == '\'' {
|
||||
if i+1 < len(card) && card[i+1] == '\'' {
|
||||
b.WriteByte('\'')
|
||||
i += 2
|
||||
continue
|
||||
}
|
||||
return b.String(), nil
|
||||
}
|
||||
b.WriteByte(card[i])
|
||||
i++
|
||||
}
|
||||
return "", base.Errf("CONTINUE card %q has an unterminated string", card[:min(20, len(card))])
|
||||
}
|
||||
|
||||
// fitsCheckKeyword validates a header keyword the caller supplies.
|
||||
// The format's commentary keywords COMMENT and HISTORY carry no
|
||||
// value, so they are refused like the structural ones.
|
||||
func fitsCheckKeyword(kw string) error {
|
||||
if len(kw) < 1 || len(kw) > 8 {
|
||||
return base.Errf("keyword %q must be 1 to 8 characters", kw)
|
||||
}
|
||||
for _, r := range kw {
|
||||
switch {
|
||||
case r >= 'A' && r <= 'Z', r >= '0' && r <= '9', r == '-', r == '_':
|
||||
default:
|
||||
return base.Errf("keyword %q may only contain A-Z, 0-9, '-' and '_'", kw)
|
||||
}
|
||||
}
|
||||
switch kw {
|
||||
case "SIMPLE", "BITPIX", "NAXIS", "EXTEND", "END", "COMMENT", "HISTORY":
|
||||
return base.Errf("keyword %q is reserved by the format", kw)
|
||||
case "BSCALE", "BZERO":
|
||||
// SaveFITS writes physical values directly; a user scaling card
|
||||
// would make conforming readers scale them a second time.
|
||||
return base.Errf("keyword %q is reserved by the format (values are written unscaled)", kw)
|
||||
}
|
||||
if rest, ok := strings.CutPrefix(kw, "NAXIS"); ok && rest != "" {
|
||||
if _, err := strconv.Atoi(rest); err == nil {
|
||||
return base.Errf("keyword %q is reserved by the format", kw)
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// fitsAxisKeyword reports whether key is a structural NAXISn card: the
|
||||
// prefix followed by a number, the same rule fitsCheckKeyword applies
|
||||
// when it refuses a reserved keyword on write. Any other spelling of
|
||||
// the prefix (NAXISREF and the like) is a user keyword, and counting it
|
||||
// as an axis used to make the reader answer "NAXIS = 1 with 2 NAXISn
|
||||
// cards" for a file that carries one.
|
||||
func fitsAxisKeyword(key string) bool {
|
||||
rest, ok := strings.CutPrefix(key, "NAXIS")
|
||||
if !ok {
|
||||
return false
|
||||
}
|
||||
_, err := strconv.Atoi(rest)
|
||||
return err == nil
|
||||
}
|
||||
|
||||
// fitsLookupN reads headers[prefix+index] without building the key on
|
||||
// the heap: the digits are formatted into a stack buffer, and the map
|
||||
// lookup over that buffer compiles to a lookup over the bytes with no
|
||||
// string conversion.
|
||||
func fitsLookupN(headers map[string]string, prefix string, index int) string {
|
||||
var kb [24]byte
|
||||
b := append(kb[:0], prefix...)
|
||||
b = strconv.AppendInt(b, int64(index), 10)
|
||||
return headers[string(b)]
|
||||
}
|
||||
|
||||
// fitsIntCard renders an integer-valued card with the value
|
||||
// right-justified in columns 11 to 30.
|
||||
func fitsIntCard(keyword string, v int) string {
|
||||
return fitsPadCard(fmt.Sprintf("%-8s= %20d", keyword, v))
|
||||
}
|
||||
|
||||
// fitsBoolCard renders a logical-valued card.
|
||||
func fitsBoolCard(keyword string, v bool) string {
|
||||
t := "F"
|
||||
if v {
|
||||
t = "T"
|
||||
}
|
||||
return fitsPadCard(fmt.Sprintf("%-8s= %20s", keyword, t))
|
||||
}
|
||||
|
||||
// fitsStringCard renders a string-valued card; the format pads the
|
||||
// quoted value to at least eight characters.
|
||||
func fitsStringCard(keyword, v string) (string, error) {
|
||||
escaped := strings.ReplaceAll(v, "'", "''")
|
||||
inner := escaped
|
||||
if len(inner) < 8 {
|
||||
inner += strings.Repeat(" ", 8-len(inner))
|
||||
}
|
||||
body := fmt.Sprintf("%-8s= '%s'", keyword, inner)
|
||||
if len(body) > 80 {
|
||||
return "", base.Errf("value for %q does not fit a card after quote escaping (%d characters)",
|
||||
keyword, len(escaped))
|
||||
}
|
||||
return fitsPadCard(body), nil
|
||||
}
|
||||
|
||||
// fitsEndCard renders the header terminator.
|
||||
func fitsEndCard() string {
|
||||
return fitsPadCard("END")
|
||||
}
|
||||
|
||||
// fitsPadCard right-pads a card body with spaces to the full 80 bytes.
|
||||
func fitsPadCard(body string) string {
|
||||
return body + strings.Repeat(" ", 80-len(body))
|
||||
}
|
||||
|
||||
// fitsUserCards renders the caller's header entries as cards, keywords
|
||||
// uppercased and sorted so the output is deterministic.
|
||||
func fitsUserCards(headers map[string]string) ([]string, error) {
|
||||
cards := make([]string, 0, len(headers))
|
||||
for _, key := range slices.Sorted(maps.Keys(headers)) {
|
||||
kw := strings.ToUpper(key)
|
||||
if err := fitsCheckKeyword(kw); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
card, err := fitsStringCard(kw, headers[key])
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
cards = append(cards, card)
|
||||
}
|
||||
return cards, nil
|
||||
}
|
||||
|
||||
// fitsCard is one parsed header card: its keyword and the value field
|
||||
// with quoting resolved.
|
||||
type fitsCard struct {
|
||||
key, value string
|
||||
}
|
||||
|
||||
// scanFITSCards walks the 80-byte cards of one header starting at off
|
||||
// and returns every value card up to the END terminator together with
|
||||
// the block-aligned header length in bytes, which is where the data
|
||||
// block begins. Blank, COMMENT and HISTORY cards carry no value and
|
||||
// are skipped before the value-indicator check, because a commentary
|
||||
// card may legitimately carry an "= " sequence in columns 9-10; a
|
||||
// header that never terminates is an error.
|
||||
func scanFITSCards(data []byte, off int) ([]fitsCard, int, error) {
|
||||
var cards []fitsCard
|
||||
seen := 0
|
||||
for pos := off; pos+80 <= len(data); pos += 80 {
|
||||
line := string(data[pos : pos+80])
|
||||
seen++
|
||||
key := strings.TrimRight(line[:8], " ")
|
||||
if key == "END" {
|
||||
return cards, (seen + fitsCardsPerBlock - 1) / fitsCardsPerBlock * 2880, nil
|
||||
}
|
||||
// Commentary cards have no value regardless of what follows
|
||||
// columns 9-10; check them before the "= " indicator.
|
||||
if key == "" || key == "COMMENT" || key == "HISTORY" {
|
||||
continue
|
||||
}
|
||||
if key == "CONTINUE" {
|
||||
// A CONTINUE card only makes sense after an open long
|
||||
// string; reaching one here means the base card never
|
||||
// ended in '&', and silently dropping it would lose the
|
||||
// caller's value.
|
||||
return nil, 0, base.Errf("CONTINUE card without an open long string")
|
||||
}
|
||||
if line[8:10] != "= " {
|
||||
continue
|
||||
}
|
||||
value, wasString, err := fitsValue(line)
|
||||
if err != nil {
|
||||
return nil, 0, err
|
||||
}
|
||||
// Long-string convention: a string value ending in '&' is
|
||||
// continued by the following CONTINUE cards, each carrying the
|
||||
// next quoted segment. Dropping them would truncate the value
|
||||
// at the card boundary.
|
||||
for wasString && strings.HasSuffix(value, "&") && pos+160 <= len(data) {
|
||||
next := string(data[pos+80 : pos+160])
|
||||
if strings.TrimSpace(next[:8]) != "CONTINUE" {
|
||||
break
|
||||
}
|
||||
pos += 80
|
||||
seen++
|
||||
cont, cerr := fitsContinueString(next)
|
||||
if cerr != nil {
|
||||
return nil, 0, cerr
|
||||
}
|
||||
value = strings.TrimSuffix(value, "&") + strings.TrimRight(cont, " ")
|
||||
}
|
||||
cards = append(cards, fitsCard{key, value})
|
||||
}
|
||||
return nil, 0, base.Errf("no END card terminates the header")
|
||||
}
|
||||
|
||||
// fitsBlockSize returns n rounded up to a whole FITS block.
|
||||
func fitsBlockSize(n int) int {
|
||||
return (n + 2879) / 2880 * 2880
|
||||
}
|
||||
|
||||
// fitsAppendCards appends the rendered cards as one header block,
|
||||
// blank-padded to the block boundary.
|
||||
func fitsAppendCards(dst []byte, cards []string) []byte {
|
||||
for _, card := range cards {
|
||||
dst = append(dst, card...)
|
||||
}
|
||||
// Blank cards pad the header to the block boundary.
|
||||
if rem := len(dst) % 2880; rem != 0 {
|
||||
blank := fitsPadCard("")
|
||||
for i := 0; i < 2880-rem; i += 80 {
|
||||
dst = append(dst, blank...)
|
||||
}
|
||||
}
|
||||
return dst
|
||||
}
|
||||
|
||||
// fitsAppendZeroPad appends zero bytes up to the block boundary.
|
||||
func fitsAppendZeroPad(dst []byte) []byte {
|
||||
if rem := len(dst) % 2880; rem != 0 {
|
||||
dst = append(dst, make([]byte, 2880-rem)...)
|
||||
}
|
||||
return dst
|
||||
}
|
||||
@@ -0,0 +1,272 @@
|
||||
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
|
||||
// SPDX-License-Identifier: MIT
|
||||
|
||||
package io
|
||||
|
||||
import (
|
||||
"encoding/binary"
|
||||
"strings"
|
||||
"testing"
|
||||
)
|
||||
|
||||
// Regression pins for hostile FITS input: negative and missing axis
|
||||
// cards, TBCOL and row-byte arithmetic that wraps, repeat counts past
|
||||
// the int range, unbounded HDU skips and garbage PCOUNT, all refused by
|
||||
// name before any payload is read.
|
||||
|
||||
// tableHostile builds a FITS table file from card bodies plus one
|
||||
// 2880-byte data block of the given row content.
|
||||
func tableHostile(cards []string, rows int) []byte {
|
||||
var b []byte
|
||||
for _, body := range cards {
|
||||
b = append(b, card(body)...)
|
||||
}
|
||||
b = append(b, card("END")...)
|
||||
if pad := len(b) % 2880; pad != 0 {
|
||||
b = append(b, make([]byte, 2880-pad)...)
|
||||
}
|
||||
b = append(b, make([]byte, 2880)...)
|
||||
_ = rows
|
||||
return b
|
||||
}
|
||||
|
||||
// TestFITSNegativeAxisCard pins the refusal of a negative or
|
||||
// missing NAXIS before the axis slice is allocated, and the image skip
|
||||
// past a zero axis in the table reader.
|
||||
func TestFITSNegativeAxisCard(t *testing.T) {
|
||||
build := func(bodies []string) []byte {
|
||||
var b []byte
|
||||
for _, body := range bodies {
|
||||
b = append(b, card(body)...)
|
||||
}
|
||||
b = append(b, card("END")...)
|
||||
if pad := len(b) % 2880; pad != 0 {
|
||||
b = append(b, make([]byte, 2880-pad)...)
|
||||
}
|
||||
return b
|
||||
}
|
||||
noNaxis := writeHostile(t, "no-naxis.fits", build([]string{
|
||||
"SIMPLE = T",
|
||||
"BITPIX = -64",
|
||||
}))
|
||||
if _, _, err := LoadFITS(noNaxis); err == nil {
|
||||
t.Fatal("expected an error for a missing NAXIS card")
|
||||
}
|
||||
negative := writeHostile(t, "neg-naxis.fits", build([]string{
|
||||
"SIMPLE = T",
|
||||
"BITPIX = -64",
|
||||
"NAXIS = -3",
|
||||
}))
|
||||
if _, _, err := LoadFITS(negative); err == nil {
|
||||
t.Fatal("expected an error for a negative NAXIS")
|
||||
}
|
||||
// A zero axis followed by another divides the running product in
|
||||
// the table reader's image skip; the skip must pass it without
|
||||
// dividing by the zeroed product and report that no table follows.
|
||||
zeroAxis := writeHostile(t, "zero-axis.fits", build([]string{
|
||||
"SIMPLE = T",
|
||||
"BITPIX = -64",
|
||||
"NAXIS = 2",
|
||||
"NAXIS1 = 0",
|
||||
"NAXIS2 = 1",
|
||||
}))
|
||||
if _, err := LoadFITSTable(zeroAxis); err == nil || !strings.Contains(err.Error(), "no table extension") {
|
||||
t.Fatalf("LoadFITSTable past a zero axis: err = %v, want the no-table report", err)
|
||||
}
|
||||
if _, _, err := LoadFITS(zeroAxis); err == nil {
|
||||
t.Fatal("LoadFITS: expected an error for a zero axis under positive NAXIS")
|
||||
}
|
||||
}
|
||||
|
||||
// TestFITSHugeNaxisRefused pins that a NAXIS beyond the NAXISn cards
|
||||
// the file carries is refused before the axis slice is allocated: a
|
||||
// NAXIS of MaxInt64 used to reach make([]int, naxis) and panic with
|
||||
// makeslice instead of answering the card-count error.
|
||||
func TestFITSHugeNaxisRefused(t *testing.T) {
|
||||
build := func(bodies []string) []byte {
|
||||
var b []byte
|
||||
for _, body := range bodies {
|
||||
b = append(b, card(body)...)
|
||||
}
|
||||
b = append(b, card("END")...)
|
||||
if pad := len(b) % 2880; pad != 0 {
|
||||
b = append(b, make([]byte, 2880-pad)...)
|
||||
}
|
||||
return b
|
||||
}
|
||||
huge := writeHostile(t, "huge-naxis.fits", build([]string{
|
||||
"SIMPLE = T",
|
||||
"BITPIX = -64",
|
||||
"NAXIS = 9223372036854775807",
|
||||
}))
|
||||
if _, _, err := LoadFITS(huge); err == nil || !strings.Contains(err.Error(), "NAXISn cards") {
|
||||
t.Fatalf("LoadFITS with NAXIS = MaxInt64: err = %v, want the card-count refusal", err)
|
||||
}
|
||||
}
|
||||
|
||||
// TestASCIITableTBCOLWrap pins the overflow-free bound on
|
||||
// TBCOL + width: a TBCOL near MaxInt64 used to wrap the sum negative,
|
||||
// pass the guard and panic on the slice.
|
||||
func TestASCIITableTBCOLWrap(t *testing.T) {
|
||||
data := tableHostile([]string{
|
||||
"XTENSION= 'TABLE '",
|
||||
"BITPIX = 8",
|
||||
"NAXIS = 2",
|
||||
"NAXIS1 = 20",
|
||||
"NAXIS2 = 1",
|
||||
"TFIELDS = 1",
|
||||
"TFORM1 = 'I10 '",
|
||||
"TBCOL1 = '9223372036854775807'",
|
||||
}, 1)
|
||||
path := writeHostile(t, "tbcol-wrap.fits", data)
|
||||
if _, err := LoadFITSTable(path); err == nil {
|
||||
t.Fatal("expected an error for a TBCOL whose sum with the width wraps")
|
||||
}
|
||||
}
|
||||
|
||||
// TestBINTABLERowBytesWrap pins the per-column bound on the
|
||||
// row prefix sum: two A columns of 2^62 repeat each used to wrap
|
||||
// rowBytes negative, skip the truncation guard and read past the file.
|
||||
func TestBINTABLERowBytesWrap(t *testing.T) {
|
||||
data := tableHostile([]string{
|
||||
"XTENSION= 'BINTABLE'",
|
||||
"BITPIX = 8",
|
||||
"NAXIS = 2",
|
||||
"NAXIS1 = 16",
|
||||
"NAXIS2 = 1",
|
||||
"TFIELDS = 2",
|
||||
"TFORM1 = '4611686018427387904A'",
|
||||
"TFORM2 = '4611686018427387904A'",
|
||||
}, 1)
|
||||
path := writeHostile(t, "rowbytes-wrap.fits", data)
|
||||
if _, err := LoadFITSTable(path); err == nil {
|
||||
t.Fatal("expected an error for column widths whose sum wraps")
|
||||
}
|
||||
}
|
||||
|
||||
// TestTFORMRepeatOverflowRefused pins the named refusal for a
|
||||
// repeat count that does not fit an int, which used to fall back to a
|
||||
// silent scalar of width 1.
|
||||
func TestTFORMRepeatOverflowRefused(t *testing.T) {
|
||||
if _, _, err := parseTFORM("99999999999999999999E"); err == nil {
|
||||
t.Fatal("parseTFORM: expected an error for an overflowing repeat")
|
||||
}
|
||||
if r, code, err := parseTFORM("16A"); err != nil || r != 16 || code != "A" {
|
||||
t.Fatalf("parseTFORM(16A) = %d, %q, %v", r, code, err)
|
||||
}
|
||||
}
|
||||
|
||||
// TestImageHDUSkipBounded pins the image-HDU skip arithmetic:
|
||||
// axis products that would wrap the int cannot skip into attacker
|
||||
// bytes; the file is refused as truncated instead.
|
||||
func TestImageHDUSkipBounded(t *testing.T) {
|
||||
var b []byte
|
||||
for _, body := range []string{
|
||||
"SIMPLE = T",
|
||||
"BITPIX = -64",
|
||||
"NAXIS = 2",
|
||||
"NAXIS1 = 2147483647",
|
||||
"NAXIS2 = 4",
|
||||
"PCOUNT = 0",
|
||||
"GCOUNT = 2",
|
||||
} {
|
||||
b = append(b, card(body)...)
|
||||
}
|
||||
b = append(b, card("END")...)
|
||||
if pad := len(b) % 2880; pad != 0 {
|
||||
b = append(b, make([]byte, 2880-pad)...)
|
||||
}
|
||||
b = append(b, make([]byte, 2880)...)
|
||||
path := writeHostile(t, "skip-wrap.fits", b)
|
||||
// The table reader skips the image HDU to look for a table after
|
||||
// it; the wrapped product used to land the skip inside the image
|
||||
// data. Now the skip is bounded and the file has no table.
|
||||
if _, err := LoadFITSTable(path); err == nil {
|
||||
t.Fatal("expected an error: no table HDU follows the bounded skip")
|
||||
}
|
||||
}
|
||||
|
||||
// TestPCOUNTGarbageRefused pins that a non-integer PCOUNT is
|
||||
// an error, not a silent zero.
|
||||
func TestPCOUNTGarbageRefused(t *testing.T) {
|
||||
var b []byte
|
||||
for _, body := range []string{
|
||||
"SIMPLE = T",
|
||||
"BITPIX = -64",
|
||||
"NAXIS = 2",
|
||||
"NAXIS1 = 2",
|
||||
"NAXIS2 = 2",
|
||||
"PCOUNT = '1e3'",
|
||||
} {
|
||||
b = append(b, card(body)...)
|
||||
}
|
||||
b = append(b, card("END")...)
|
||||
if pad := len(b) % 2880; pad != 0 {
|
||||
b = append(b, make([]byte, 2880-pad)...)
|
||||
}
|
||||
b = append(b, make([]byte, 2880)...)
|
||||
path := writeHostile(t, "pcount.fits", b)
|
||||
if _, err := LoadFITSTable(path); err == nil {
|
||||
t.Fatal("expected an error for a non-integer PCOUNT")
|
||||
}
|
||||
}
|
||||
|
||||
// TestFITSNAXISnOrder pins that NAXISn cards are bound by
|
||||
// their axis number: a header writing NAXIS2 before NAXIS1 under
|
||||
// NAXIS = 2 with matching extents stays the image it declares, and a
|
||||
// missing axis is an error.
|
||||
func TestFITSNAXISnOrder(t *testing.T) {
|
||||
build := func(bodies []string) []byte {
|
||||
var b []byte
|
||||
for _, body := range bodies {
|
||||
b = append(b, card(body)...)
|
||||
}
|
||||
b = append(b, card("END")...)
|
||||
if pad := len(b) % 2880; pad != 0 {
|
||||
b = append(b, make([]byte, 2880-pad)...)
|
||||
}
|
||||
// 2x2 float64 image data.
|
||||
for range 4 {
|
||||
b = binary.BigEndian.AppendUint64(b, 1)
|
||||
}
|
||||
return b
|
||||
}
|
||||
// Reversed card order: the image is still 3 by 2.
|
||||
path := writeHostile(t, "order.fits", build([]string{
|
||||
"SIMPLE = T",
|
||||
"BITPIX = -64",
|
||||
"NAXIS = 2",
|
||||
"NAXIS2 = 2",
|
||||
"NAXIS1 = 3",
|
||||
}))
|
||||
// 3x2 = 6 doubles but only 4 present: truncated either way, so the
|
||||
// claim is checked by the value the reader reports. Give it enough
|
||||
// data instead: rebuild with full payload.
|
||||
full := build([]string{
|
||||
"SIMPLE = T",
|
||||
"BITPIX = -64",
|
||||
"NAXIS = 2",
|
||||
"NAXIS2 = 2",
|
||||
"NAXIS1 = 3",
|
||||
})
|
||||
full = append(full, make([]byte, 2880)...)
|
||||
path = writeHostile(t, "order-full.fits", full)
|
||||
a, _, err := LoadFITS(path)
|
||||
if err != nil {
|
||||
t.Fatalf("LoadFITS with reversed NAXISn cards: %v", err)
|
||||
}
|
||||
if a.Shape()[0] != 2 || a.Shape()[1] != 3 {
|
||||
t.Fatalf("shape = %v, want [2 3] (Fortran order, NAXIS1 = 3)", a.Shape())
|
||||
}
|
||||
// A missing axis is refused by name.
|
||||
missing := build([]string{
|
||||
"SIMPLE = T",
|
||||
"BITPIX = -64",
|
||||
"NAXIS = 2",
|
||||
"NAXIS1 = 2",
|
||||
})
|
||||
mpath := writeHostile(t, "missing.fits", missing)
|
||||
if _, _, err := LoadFITS(mpath); err == nil || !strings.Contains(err.Error(), "NAXIS") {
|
||||
t.Fatalf("err = %v, want the missing NAXISn card named", err)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,206 @@
|
||||
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
|
||||
// SPDX-License-Identifier: MIT
|
||||
|
||||
package io
|
||||
|
||||
import (
|
||||
"encoding/binary"
|
||||
"fmt"
|
||||
"math"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"sourcedock.dev/petrbalvin/tensor/internal/core"
|
||||
)
|
||||
|
||||
// Regression pins for the FITS table codecs: the binary and ASCII
|
||||
// column round trips, the header cards that size them, and the hostile
|
||||
// counts and widths the reader must refuse by name.
|
||||
|
||||
// card renders one 80-byte FITS header card from a full-line body.
|
||||
func card(body string) []byte {
|
||||
return []byte(body + strings.Repeat(" ", 80-len(body)))
|
||||
}
|
||||
|
||||
// cardBlock concatenates cards and pads to the FITS block size.
|
||||
func cardBlock(cards ...[]byte) []byte {
|
||||
var b []byte
|
||||
for _, c := range cards {
|
||||
b = append(b, c...)
|
||||
}
|
||||
if rem := len(b) % 2880; rem != 0 {
|
||||
b = append(b, make([]byte, 2880-rem)...)
|
||||
}
|
||||
return b
|
||||
}
|
||||
|
||||
// TestBinaryStringColumnRoundTrip pins the nA binary column contract:
|
||||
// the reader must accept the repeat-count width the writer emits, so
|
||||
// SaveFITSTable to LoadFITSTable survives character columns.
|
||||
func TestBinaryStringColumnRoundTrip(t *testing.T) {
|
||||
names := []string{"M31", "NGC 1275", "SMC"}
|
||||
cols := []FITSTableColumn{
|
||||
{Name: "NAME", Form: "8A", Text: names},
|
||||
{Name: "FLUX", Form: "D", Data: mustFloats(t, []float64{1, 2, 3}, 3)},
|
||||
}
|
||||
path := filepath.Join(t.TempDir(), "names.fits")
|
||||
if err := SaveFITSTable(path, false, cols, nil); err != nil {
|
||||
t.Fatalf("SaveFITSTable: %v", err)
|
||||
}
|
||||
table, err := LoadFITSTable(path)
|
||||
if err != nil {
|
||||
t.Fatalf("LoadFITSTable: %v", err)
|
||||
}
|
||||
for i, want := range names {
|
||||
if table.Text[0][i] != want {
|
||||
t.Fatalf("NAME[%d] = %q, want %q", i, table.Text[0][i], want)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// TestTableTSCALTZERO pins the scaling keywords: a binary int column
|
||||
// with TSCAL/TZERO must read back through the affine map, the standard
|
||||
// unsigned-integer convention included.
|
||||
func TestTableTSCALTZERO(t *testing.T) {
|
||||
raw := []int64{1, 2, 3}
|
||||
ints := core.New(core.Int, len(raw))
|
||||
copy(ints.RawInts(), raw)
|
||||
cols := []FITSTableColumn{
|
||||
{Name: "U16", Form: "K", Data: ints},
|
||||
}
|
||||
path := filepath.Join(t.TempDir(), "scaled.fits")
|
||||
headers := map[string]string{"TSCAL1": "1", "TZERO1": "32768"}
|
||||
if err := SaveFITSTable(path, false, cols, headers); err != nil {
|
||||
t.Fatalf("SaveFITSTable: %v", err)
|
||||
}
|
||||
table, err := LoadFITSTable(path)
|
||||
if err != nil {
|
||||
t.Fatalf("LoadFITSTable: %v", err)
|
||||
}
|
||||
col := table.Columns[0]
|
||||
for i, v := range raw {
|
||||
want := float64(v) + 32768
|
||||
got := col.FloatAt(i)
|
||||
if got != want {
|
||||
t.Fatalf("U16[%d] = %g, want %g", i, got, want)
|
||||
}
|
||||
}
|
||||
// A fractional scale promotes the column to float.
|
||||
if err := SaveFITSTable(path, false, cols, map[string]string{"TSCAL1": "0.5", "TZERO1": "0"}); err != nil {
|
||||
t.Fatalf("SaveFITSTable: %v", err)
|
||||
}
|
||||
table, err = LoadFITSTable(path)
|
||||
if err != nil {
|
||||
t.Fatalf("LoadFITSTable: %v", err)
|
||||
}
|
||||
if table.Columns[0].Dtype() != core.Float {
|
||||
t.Fatalf("scaled int column dtype %s, want float", table.Columns[0].Dtype())
|
||||
}
|
||||
}
|
||||
|
||||
// TestLoadFITSTableSkips16BitImage pins the image-HDU skip arithmetic
|
||||
// on the case that used to break: a positive BITPIX (16) with a
|
||||
// non-square shape, where |BITPIX|/8 · Π NAXISn ≠ Σ NAXISn.
|
||||
func TestLoadFITSTableSkips16BitImage(t *testing.T) {
|
||||
// A 5×3 16-bit image (NAXIS1 = 5 is the fastest axis): 15 pixels,
|
||||
// 30 bytes of data.
|
||||
hdr := cardBlock(
|
||||
card("SIMPLE = T"),
|
||||
card("BITPIX = 16"),
|
||||
card("NAXIS = 2"),
|
||||
card("NAXIS1 = 5"),
|
||||
card("NAXIS2 = 3"),
|
||||
card("END"),
|
||||
)
|
||||
data := make([]byte, 0, 30)
|
||||
for range 15 {
|
||||
data = binary.BigEndian.AppendUint16(data, 42)
|
||||
}
|
||||
image := append(hdr, data...)
|
||||
if rem := len(image) % 2880; rem != 0 {
|
||||
image = append(image, make([]byte, 2880-rem)...)
|
||||
}
|
||||
|
||||
cols := []FITSTableColumn{
|
||||
{Name: "V", Form: "K", Data: mustInts(t, []int64{7, 8}, 2)},
|
||||
}
|
||||
tablePath := filepath.Join(t.TempDir(), "table.fits")
|
||||
if err := SaveFITSTable(tablePath, false, cols, nil); err != nil {
|
||||
t.Fatalf("SaveFITSTable: %v", err)
|
||||
}
|
||||
tableBytes, err := osReadFile(tablePath)
|
||||
if err != nil {
|
||||
t.Fatalf("read table: %v", err)
|
||||
}
|
||||
// Drop the table file's empty primary HDU (first block), keep the
|
||||
// extension behind the hand-built image.
|
||||
path := filepath.Join(t.TempDir(), "im_then_table.fits")
|
||||
if err := osWriteFile(path, image, tableBytes[2880:]); err != nil {
|
||||
t.Fatalf("write: %v", err)
|
||||
}
|
||||
table, err := LoadFITSTable(path)
|
||||
if err != nil {
|
||||
t.Fatalf("LoadFITSTable after a 16-bit image: %v", err)
|
||||
}
|
||||
if table.Rows != 2 || table.Columns[0].RawInts()[1] != 8 {
|
||||
t.Fatalf("table landed wrong: rows %d, V[1] = %d", table.Rows, table.Columns[0].RawInts()[1])
|
||||
}
|
||||
}
|
||||
|
||||
// TestLoadFITSBSCALE pins image scaling: physical = raw·BSCALE + BZERO.
|
||||
func TestLoadFITSBSCALE(t *testing.T) {
|
||||
hdr := cardBlock(
|
||||
card("SIMPLE = T"),
|
||||
card("BITPIX = -64"),
|
||||
card("NAXIS = 1"),
|
||||
card("NAXIS1 = 2"),
|
||||
card("BSCALE = 2.0"),
|
||||
card("BZERO = 10.0"),
|
||||
card("END"),
|
||||
)
|
||||
var data []byte
|
||||
data = binary.BigEndian.AppendUint64(data, math.Float64bits(1))
|
||||
data = binary.BigEndian.AppendUint64(data, math.Float64bits(2))
|
||||
path := filepath.Join(t.TempDir(), "scaled.fits")
|
||||
var payload []byte
|
||||
payload = append(payload, data...)
|
||||
if rem := len(payload) % 2880; rem != 0 {
|
||||
payload = append(payload, make([]byte, 2880-rem)...)
|
||||
}
|
||||
if err := osWriteFile(path, hdr, payload); err != nil {
|
||||
t.Fatalf("write: %v", err)
|
||||
}
|
||||
img, _, err := LoadFITS(path)
|
||||
if err != nil {
|
||||
t.Fatalf("LoadFITS: %v", err)
|
||||
}
|
||||
if img.FloatAt(0) != 12 || img.FloatAt(1) != 14 {
|
||||
t.Fatalf("scaled pixels (%g, %g), want (12, 14)", img.FloatAt(0), img.FloatAt(1))
|
||||
}
|
||||
}
|
||||
|
||||
// TestFITSContinueLongString pins the long-string convention: a value
|
||||
// ending in '&' continues on the CONTINUE cards instead of truncating.
|
||||
func TestFITSContinueLongString(t *testing.T) {
|
||||
hdr := cardBlock(
|
||||
card("SIMPLE = T"),
|
||||
card("BITPIX = 8"),
|
||||
card("NAXIS = 0"),
|
||||
card(fmt.Sprintf("%-8s= '%s&'", "LONGKEY", "the first part of a long value ")),
|
||||
card(fmt.Sprintf("CONTINUE '%s'", "and the second part")),
|
||||
card("END"),
|
||||
)
|
||||
path := filepath.Join(t.TempDir(), "long.fits")
|
||||
if err := osWriteFile(path, hdr); err != nil {
|
||||
t.Fatalf("write: %v", err)
|
||||
}
|
||||
_, headers, err := LoadFITS(path)
|
||||
if err != nil {
|
||||
t.Fatalf("LoadFITS: %v", err)
|
||||
}
|
||||
want := "the first part of a long value and the second part"
|
||||
if headers["LONGKEY"] != want {
|
||||
t.Fatalf("LONGKEY = %q, want %q", headers["LONGKEY"], want)
|
||||
}
|
||||
}
|
||||
+338
@@ -0,0 +1,338 @@
|
||||
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
|
||||
// SPDX-License-Identifier: MIT
|
||||
|
||||
package io
|
||||
|
||||
import (
|
||||
"encoding/binary"
|
||||
"math"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"sourcedock.dev/petrbalvin/tensor/internal/core"
|
||||
"strings"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func fitsTempPath(t *testing.T) string {
|
||||
t.Helper()
|
||||
return filepath.Join(t.TempDir(), "image.fits")
|
||||
}
|
||||
|
||||
// padCard right-pads a card body with spaces to the full 80 bytes.
|
||||
func padCard(text string) []byte {
|
||||
return []byte(text + strings.Repeat(" ", 80-len(text)))
|
||||
}
|
||||
|
||||
// TestFITSRoundTripFloat64 moves a rank-2 float64 image with negative
|
||||
// and non-round values through the file format and back, values and
|
||||
// header strings included.
|
||||
func TestFITSRoundTripFloat64(t *testing.T) {
|
||||
a := mustFloats(t, []float64{
|
||||
1.5, -2.25, 3.125, 4,
|
||||
-5.5, 6.75, -7.875, 8,
|
||||
9.25, -10.5, 11.125, -12,
|
||||
}, 3, 4)
|
||||
path := fitsTempPath(t)
|
||||
headers := map[string]string{
|
||||
"OBJECT": "M31",
|
||||
"OBSERVER": "petr's dome",
|
||||
"EXPTIME": "600",
|
||||
}
|
||||
if err := SaveFITS(path, a, headers); err != nil {
|
||||
t.Fatalf("SaveFITS: %v", err)
|
||||
}
|
||||
back, hdr, err := LoadFITS(path)
|
||||
if err != nil {
|
||||
t.Fatalf("LoadFITS: %v", err)
|
||||
}
|
||||
if back.Dtype() != core.Float || back.NDim() != 2 || back.Shape()[0] != 3 || back.Shape()[1] != 4 {
|
||||
t.Fatalf("shape/dtype mismatch: %v %s", back.Shape(), back.Dtype())
|
||||
}
|
||||
for i := range a.Len() {
|
||||
if back.FloatAt(i) != a.FloatAt(i) {
|
||||
t.Fatalf("value[%d] = %g, want %g", i, back.FloatAt(i), a.FloatAt(i))
|
||||
}
|
||||
}
|
||||
for key, want := range map[string]string{
|
||||
"OBJECT": "M31",
|
||||
"OBSERVER": "petr's dome",
|
||||
"EXPTIME": "600",
|
||||
} {
|
||||
if hdr[key] != want {
|
||||
t.Fatalf("header %q = %q, want %q", key, hdr[key], want)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// TestFITSAxisConvention pins the wire format against the raw bytes:
|
||||
// NAXIS1 must carry the fastest (last Go) axis and the payload must
|
||||
// be big-endian in flat row-major order.
|
||||
func TestFITSAxisConvention(t *testing.T) {
|
||||
a := mustFloats(t, []float64{1, 2, 3, 4, 5, 6}, 2, 3)
|
||||
path := fitsTempPath(t)
|
||||
if err := SaveFITS(path, a, nil); err != nil {
|
||||
t.Fatalf("SaveFITS: %v", err)
|
||||
}
|
||||
raw, err := os.ReadFile(path)
|
||||
if err != nil {
|
||||
t.Fatalf("ReadFile: %v", err)
|
||||
}
|
||||
if len(raw)%2880 != 0 {
|
||||
t.Fatalf("file length %d is not a multiple of 2880", len(raw))
|
||||
}
|
||||
header := raw[:2880]
|
||||
for _, want := range []string{
|
||||
"SIMPLE = T",
|
||||
"BITPIX = -64",
|
||||
"NAXIS = 2",
|
||||
"NAXIS1 = 3",
|
||||
"NAXIS2 = 2",
|
||||
} {
|
||||
if !strings.Contains(string(header), want) {
|
||||
t.Fatalf("header misses %q", want)
|
||||
}
|
||||
}
|
||||
endAt := strings.Index(string(header), "END")
|
||||
if endAt < 0 {
|
||||
t.Fatal("header has no END card")
|
||||
}
|
||||
var word [8]byte
|
||||
for i := range 6 {
|
||||
binary.BigEndian.PutUint64(word[:], math.Float64bits(float64(i+1)))
|
||||
got := raw[2880+i*8 : 2880+i*8+8]
|
||||
if string(got) != string(word[:]) {
|
||||
t.Fatalf("payload word %d = % x, want % x", i, got, word)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// TestFITSRoundTripFloat32 keeps the float32 element type through the
|
||||
// round trip, which is what BITPIX −32 stores.
|
||||
func TestFITSRoundTripFloat32(t *testing.T) {
|
||||
a := core.New(core.Float32, 5)
|
||||
for i := range 5 {
|
||||
a.RawFloat32s()[i] = float32(i) * 1.25
|
||||
}
|
||||
path := fitsTempPath(t)
|
||||
if err := SaveFITS(path, a, nil); err != nil {
|
||||
t.Fatalf("SaveFITS: %v", err)
|
||||
}
|
||||
back, _, err := LoadFITS(path)
|
||||
if err != nil {
|
||||
t.Fatalf("LoadFITS: %v", err)
|
||||
}
|
||||
if back.Dtype() != core.Float32 {
|
||||
t.Fatalf("dtype = %s, want core.Float32", back.Dtype())
|
||||
}
|
||||
for i := range 5 {
|
||||
if back.RawFloat32s()[i] != a.RawFloat32s()[i] {
|
||||
t.Fatalf("value[%d] = %g, want %g", i, back.RawFloat32s()[i], a.RawFloat32s()[i])
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// TestFITSSkipsValuelessCards checks COMMENT and HISTORY cards are
|
||||
// tolerated and skipped rather than parsed as values.
|
||||
func TestFITSSkipsValuelessCards(t *testing.T) {
|
||||
a := mustFloats(t, []float64{1, 2, 3, 4}, 2, 2)
|
||||
path := fitsTempPath(t)
|
||||
if err := SaveFITS(path, a, map[string]string{"OBJECT": "TEST"}); err != nil {
|
||||
t.Fatalf("SaveFITS: %v", err)
|
||||
}
|
||||
raw, err := os.ReadFile(path)
|
||||
if err != nil {
|
||||
t.Fatalf("ReadFile: %v", err)
|
||||
}
|
||||
// Splice two valueless cards in before the END card and re-pad the
|
||||
// header to the block boundary. The END card is matched in full
|
||||
// ("END" plus its padding) because EXTEND contains the same three
|
||||
// letters.
|
||||
endCard := strings.Index(string(raw), "END"+strings.Repeat(" ", 77))
|
||||
if endCard < 0 {
|
||||
t.Fatal("no END card")
|
||||
}
|
||||
spliced := append([]byte{}, raw[:endCard]...)
|
||||
spliced = append(spliced, padCard("COMMENT a note without a value")...)
|
||||
spliced = append(spliced, padCard("HISTORY an audit trail entry")...)
|
||||
spliced = append(spliced, padCard("END")...)
|
||||
for len(spliced)%2880 != 0 {
|
||||
spliced = append(spliced, padCard("")...)
|
||||
}
|
||||
spliced = append(spliced, raw[2880:]...)
|
||||
if err := os.WriteFile(path, spliced, 0o644); err != nil {
|
||||
t.Fatalf("WriteFile: %v", err)
|
||||
}
|
||||
back, hdr, err := LoadFITS(path)
|
||||
if err != nil {
|
||||
t.Fatalf("LoadFITS: %v", err)
|
||||
}
|
||||
if back.Len() != 4 || back.FloatAt(3) != 4 {
|
||||
t.Fatalf("payload damaged: %v %v", back.Shape(), back.RawFloats())
|
||||
}
|
||||
if hdr["OBJECT"] != "TEST" {
|
||||
t.Fatalf("OBJECT = %q, want %q", hdr["OBJECT"], "TEST")
|
||||
}
|
||||
if _, ok := hdr["COMMENT"]; ok {
|
||||
t.Fatal("COMMENT card must not enter the header map")
|
||||
}
|
||||
}
|
||||
|
||||
// TestFITSErrors covers every refusal path in the format layer.
|
||||
func TestFITSErrors(t *testing.T) {
|
||||
path := fitsTempPath(t)
|
||||
|
||||
ints, ierr := core.FromInts([]int64{1, 2, 3, 4}, 2, 2)
|
||||
if ierr != nil {
|
||||
t.Fatalf("FromInts: %v", ierr)
|
||||
}
|
||||
if err := SaveFITS(path, ints, nil); err == nil {
|
||||
t.Fatal("int64 input: want an error")
|
||||
}
|
||||
if err := SaveFITS(path, mustComplexes(t, []complex128{1, 2, 3, 4}, 2, 2), nil); err == nil {
|
||||
t.Fatal("complex input: want an error")
|
||||
}
|
||||
ok := mustFloats(t, []float64{1, 2}, 2)
|
||||
if err := SaveFITS(path, ok, map[string]string{"TOOLONGKEYWORD": "x"}); err == nil {
|
||||
t.Fatal("long keyword: want an error")
|
||||
}
|
||||
if err := SaveFITS(path, ok, map[string]string{"BITPIX": "x"}); err == nil {
|
||||
t.Fatal("reserved keyword: want an error")
|
||||
}
|
||||
if err := SaveFITS(path, ok, map[string]string{"naxis2": "x"}); err == nil {
|
||||
t.Fatal("NAXISn keyword: want an error")
|
||||
}
|
||||
if err := SaveFITS(path, ok, map[string]string{"BAD KEY": "x"}); err == nil {
|
||||
t.Fatal("space in keyword: want an error")
|
||||
}
|
||||
if err := SaveFITS(path, ok, map[string]string{"NOTE": strings.Repeat("x", 69)}); err == nil {
|
||||
t.Fatal("overlong value: want an error")
|
||||
}
|
||||
|
||||
// A valid file to damage in every way the parser must catch.
|
||||
if err := SaveFITS(path, ok, nil); err != nil {
|
||||
t.Fatalf("SaveFITS: %v", err)
|
||||
}
|
||||
raw, err := os.ReadFile(path)
|
||||
if err != nil {
|
||||
t.Fatalf("ReadFile: %v", err)
|
||||
}
|
||||
with := func(mutate func([]byte) []byte) {
|
||||
t.Helper()
|
||||
if err := os.WriteFile(path, mutate(append([]byte{}, raw...)), 0o644); err != nil {
|
||||
t.Fatalf("WriteFile: %v", err)
|
||||
}
|
||||
if _, _, err := LoadFITS(path); err == nil {
|
||||
t.Fatal("damaged file: want an error")
|
||||
}
|
||||
}
|
||||
with(func(b []byte) []byte { return b[:100] }) // no END card
|
||||
with(func(b []byte) []byte { return b[:2880] }) // data truncated
|
||||
with(func(b []byte) []byte { b[30] = 'F'; return b }) // SIMPLE = F
|
||||
with(func(b []byte) []byte { copy(b[11:20], "XTENSION"); return b }) // wrong first card
|
||||
|
||||
// BITPIX 8 (unsigned bytes) is outside the supported image types.
|
||||
unsupported := append(append(append(append([]byte{},
|
||||
padCard("SIMPLE = T")...),
|
||||
padCard("BITPIX = 8")...),
|
||||
padCard("NAXIS = 1")...),
|
||||
padCard("NAXIS1 = 4")...)
|
||||
unsupported = append(unsupported, padCard("END")...)
|
||||
unsupported = append(unsupported, make([]byte, 2880-5*80)...)
|
||||
unsupported = append(unsupported, make([]byte, 2880)...)
|
||||
if err := os.WriteFile(path, unsupported, 0o644); err != nil {
|
||||
t.Fatalf("WriteFile: %v", err)
|
||||
}
|
||||
if _, _, err := LoadFITS(path); err == nil {
|
||||
t.Fatal("BITPIX 8: want an error")
|
||||
}
|
||||
|
||||
// A hostile header whose axis product overflows int must be
|
||||
// refused as truncated data, never panic the process.
|
||||
hostile := append(append(append(append([]byte{},
|
||||
padCard("SIMPLE = T")...),
|
||||
padCard("BITPIX = -64")...),
|
||||
padCard("NAXIS = 2")...),
|
||||
padCard("NAXIS1 = 1099511627776")...)
|
||||
hostile = append(hostile, padCard("NAXIS2 = 1099511627776")...)
|
||||
hostile = append(hostile, padCard("END")...)
|
||||
hostile = append(hostile, make([]byte, 2880-6*80)...)
|
||||
hostile = append(hostile, make([]byte, 2880)...)
|
||||
if err := os.WriteFile(path, hostile, 0o644); err != nil {
|
||||
t.Fatalf("WriteFile: %v", err)
|
||||
}
|
||||
if _, _, err := LoadFITS(path); err == nil {
|
||||
t.Fatal("overflowing axis product: want an error")
|
||||
}
|
||||
|
||||
// An XTENSION-first file is an extension, not a primary image.
|
||||
ext := append([]byte{}, padCard("XTENSION= 'IMAGE '")...)
|
||||
ext = append(ext, raw[80:]...)
|
||||
if err := os.WriteFile(path, ext, 0o644); err != nil {
|
||||
t.Fatalf("WriteFile: %v", err)
|
||||
}
|
||||
if _, _, err := LoadFITS(path); err == nil {
|
||||
t.Fatal("extension header: want an error")
|
||||
}
|
||||
|
||||
// The intact file still loads after all the damage around it.
|
||||
if err := os.WriteFile(path, raw, 0o644); err != nil {
|
||||
t.Fatalf("WriteFile: %v", err)
|
||||
}
|
||||
if _, _, err := LoadFITS(path); err != nil {
|
||||
t.Fatalf("intact file: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
// TestCommentaryCardsAreNotValueCards pins the commentary rule: a
|
||||
// COMMENT or HISTORY card may legitimately carry an "= " sequence in
|
||||
// columns 9-10, and such a card must be skipped as commentary, not
|
||||
// parsed as a keyword with a value.
|
||||
func TestCommentaryCardsAreNotValueCards(t *testing.T) {
|
||||
cards := []string{
|
||||
fitsBoolCard("SIMPLE", true),
|
||||
fitsIntCard("BITPIX", -64),
|
||||
fitsIntCard("NAXIS", 1),
|
||||
fitsIntCard("NAXIS1", 2),
|
||||
fitsPadCard("COMMENT = this looks like a value card"),
|
||||
fitsPadCard("HISTORY = so does this one"),
|
||||
fitsStringCardRaw("OBSERVER", "tester"),
|
||||
fitsEndCard(),
|
||||
}
|
||||
data := fitsAppendCards(nil, cards)
|
||||
payload := []byte{0x3f, 0xf0, 0, 0, 0, 0, 0, 0, 0x40, 0, 0, 0, 0, 0, 0, 0} // 1.0, 2.0
|
||||
data = append(data, payload...)
|
||||
data = fitsAppendZeroPad(data)
|
||||
|
||||
path := filepath.Join(t.TempDir(), "commentary.fits")
|
||||
if err := osWriteFile(path, data); err != nil {
|
||||
t.Fatalf("os.WriteFile: %v", err)
|
||||
}
|
||||
img, headers, err := LoadFITS(path)
|
||||
if err != nil {
|
||||
t.Fatalf("LoadFITS: %v", err)
|
||||
}
|
||||
if img.Len() != 2 || img.FloatAt(0) != 1 || img.FloatAt(1) != 2 {
|
||||
t.Fatalf("image = %s, want [1, 2]", img)
|
||||
}
|
||||
if _, ok := headers["COMMENT"]; ok {
|
||||
t.Error("COMMENT was parsed as a value keyword")
|
||||
}
|
||||
if _, ok := headers["HISTORY"]; ok {
|
||||
t.Error("HISTORY was parsed as a value keyword")
|
||||
}
|
||||
if headers["OBSERVER"] != "tester" {
|
||||
t.Errorf("OBSERVER = %q, want tester", headers["OBSERVER"])
|
||||
}
|
||||
}
|
||||
|
||||
// TestFitsCheckKeywordRefusesCommentary pins that COMMENT and HISTORY
|
||||
// are refused as user keywords: they carry no value in the format.
|
||||
func TestFitsCheckKeywordRefusesCommentary(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
img := mustFloats(t, []float64{1, 2}, 2)
|
||||
for _, kw := range []string{"COMMENT", "HISTORY"} {
|
||||
if err := SaveFITS(filepath.Join(dir, strings.ToLower(kw)+".fits"), img, map[string]string{kw: "x"}); err == nil {
|
||||
t.Errorf("SaveFITS accepted the reserved keyword %q", kw)
|
||||
}
|
||||
}
|
||||
}
|
||||
+1082
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,80 @@
|
||||
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
|
||||
// SPDX-License-Identifier: MIT
|
||||
|
||||
package io
|
||||
|
||||
import (
|
||||
"encoding/binary"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"testing"
|
||||
|
||||
"sourcedock.dev/petrbalvin/tensor/internal/core"
|
||||
)
|
||||
|
||||
// The binary-table reader resolves header keywords through a map whose
|
||||
// documented rule is first-occurrence-wins, the behaviour the previous
|
||||
// linear scan answered. A hostile header may repeat a structural card
|
||||
// with a different value: this pin fixes that the reader follows the
|
||||
// first value of each repeated card.
|
||||
func TestLoadFITSTableRepeatedKeywordsFirstWins(t *testing.T) {
|
||||
body := make([]byte, 16)
|
||||
for r := range 2 {
|
||||
binary.BigEndian.PutUint64(body[r*8:], uint64(1000+r))
|
||||
}
|
||||
cards := []string{
|
||||
fitsStringCardRaw("XTENSION", "BINTABLE"),
|
||||
fitsIntCard("BITPIX", 8),
|
||||
fitsIntCard("NAXIS", 2),
|
||||
fitsIntCard("NAXIS1", 8),
|
||||
fitsIntCard("NAXIS2", 2),
|
||||
fitsIntCard("NAXIS2", 99),
|
||||
fitsIntCard("PCOUNT", 0),
|
||||
fitsIntCard("GCOUNT", 1),
|
||||
fitsIntCard("TFIELDS", 1),
|
||||
fitsIntCard("TFIELDS", 5),
|
||||
fitsStringCardRaw("TTYPE1", "COL1"),
|
||||
fitsStringCardRaw("TFORM1", "K"),
|
||||
fitsStringCardRaw("TFORM1", "D"),
|
||||
fitsEndCard(),
|
||||
}
|
||||
out := fitsAppendCards(nil, []string{
|
||||
fitsBoolCard("SIMPLE", true),
|
||||
fitsIntCard("BITPIX", 8),
|
||||
fitsIntCard("NAXIS", 0),
|
||||
fitsBoolCard("EXTEND", true),
|
||||
fitsEndCard(),
|
||||
})
|
||||
out = fitsAppendCards(out, cards)
|
||||
out = append(out, body...)
|
||||
out = fitsAppendZeroPad(out)
|
||||
path := filepath.Join(t.TempDir(), "repeated.fits")
|
||||
if err := os.WriteFile(path, out, 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
table, err := LoadFITSTable(path)
|
||||
if err != nil {
|
||||
t.Fatalf("LoadFITSTable: %v", err)
|
||||
}
|
||||
if table.Rows != 2 {
|
||||
t.Fatalf("rows = %d, want 2 from the first NAXIS2", table.Rows)
|
||||
}
|
||||
if len(table.Columns) != 1 {
|
||||
t.Fatalf("columns = %d, want 1 from the first TFIELDS", len(table.Columns))
|
||||
}
|
||||
col := table.Columns[0]
|
||||
if col == nil {
|
||||
t.Fatal("the first column came back nil; the repeated TFORM1 must still decode the first form")
|
||||
}
|
||||
if len(table.Text) > 0 && table.Text[0] != nil {
|
||||
t.Fatalf("column text = %v, want nil for a K column decoded from the first TFORM1", table.Text[0])
|
||||
}
|
||||
if col.Dtype() != core.Int {
|
||||
t.Fatalf("column dtype = %s, want Int from the first TFORM1 (K)", col.Dtype())
|
||||
}
|
||||
for i, want := range []int64{1000, 1001} {
|
||||
if got := col.RawInts()[i]; got != want {
|
||||
t.Fatalf("value %d = %d, want %d", i, got, want)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,441 @@
|
||||
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
|
||||
// SPDX-License-Identifier: MIT
|
||||
|
||||
package io
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"encoding/binary"
|
||||
"math"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"sourcedock.dev/petrbalvin/tensor/internal/core"
|
||||
"testing"
|
||||
)
|
||||
|
||||
// mustFloats32 builds a float32 array.
|
||||
func mustFloats32(t *testing.T, vals []float32, shape ...int) *core.Array {
|
||||
t.Helper()
|
||||
if len(shape) == 0 {
|
||||
shape = []int{len(vals)}
|
||||
}
|
||||
a, err := core.FromFloat32s(vals, shape...)
|
||||
if err != nil {
|
||||
t.Fatalf("FromFloat32s: %v", err)
|
||||
}
|
||||
return a
|
||||
}
|
||||
|
||||
// mustInts builds an int array.
|
||||
func mustInts(t *testing.T, vals []int64, shape ...int) *core.Array {
|
||||
t.Helper()
|
||||
if len(shape) == 0 {
|
||||
shape = []int{len(vals)}
|
||||
}
|
||||
a, err := core.FromInts(vals, shape...)
|
||||
if err != nil {
|
||||
t.Fatalf("FromInts: %v", err)
|
||||
}
|
||||
return a
|
||||
}
|
||||
|
||||
// osReadFile and osWriteFile wrap the os calls so the test table
|
||||
// splicing stays terse.
|
||||
func osReadFile(path string) ([]byte, error) { return os.ReadFile(path) }
|
||||
|
||||
func osWriteFile(path string, parts ...[]byte) error {
|
||||
var out []byte
|
||||
for _, p := range parts {
|
||||
out = append(out, p...)
|
||||
}
|
||||
return os.WriteFile(path, out, 0o644)
|
||||
}
|
||||
|
||||
// starCatalogue returns a deterministic three-column table: integer
|
||||
// identifiers, float64 magnitudes and float32 temperatures.
|
||||
func starCatalogue(t *testing.T) ([]FITSTableColumn, int) {
|
||||
t.Helper()
|
||||
const n = 5
|
||||
ids := make([]float64, n)
|
||||
for i := range n {
|
||||
ids[i] = float64(1000 + i)
|
||||
}
|
||||
ints := make([]int64, n)
|
||||
for i := range n {
|
||||
ints[i] = int64(1000 + i)
|
||||
}
|
||||
mags := make([]float64, n)
|
||||
for i := range n {
|
||||
mags[i] = math.Sin(float64(3*i+1)) * 5
|
||||
}
|
||||
temps := make([]float32, n)
|
||||
for i := range n {
|
||||
temps[i] = float32(3000 + 700*i)
|
||||
}
|
||||
idArr := core.New(core.Int, n)
|
||||
copy(idArr.RawInts(), ints)
|
||||
cols := []FITSTableColumn{
|
||||
{Name: "ID", Unit: "", Form: "K", Data: idArr},
|
||||
{Name: "MAG", Unit: "mag", Form: "D", Data: mustFloats(t, mags, n)},
|
||||
{Name: "TEMP", Unit: "K", Form: "E", Data: mustFloats32(t, temps, n)},
|
||||
}
|
||||
return cols, n
|
||||
}
|
||||
|
||||
// TestSaveLoadFITSTableBinary round-trips a binary table: names,
|
||||
// units, and every value read back unchanged.
|
||||
func TestSaveLoadFITSTableBinary(t *testing.T) {
|
||||
cols, n := starCatalogue(t)
|
||||
path := filepath.Join(t.TempDir(), "catalogue.fits")
|
||||
if err := SaveFITSTable(path, false, cols, map[string]string{"ORIGIN": "tensor test"}); err != nil {
|
||||
t.Fatalf("SaveFITSTable: %v", err)
|
||||
}
|
||||
table, err := LoadFITSTable(path)
|
||||
if err != nil {
|
||||
t.Fatalf("LoadFITSTable: %v", err)
|
||||
}
|
||||
if table.Kind != "BINTABLE" {
|
||||
t.Fatalf("kind %q, want BINTABLE", table.Kind)
|
||||
}
|
||||
if table.Rows != n {
|
||||
t.Fatalf("rows = %d, want %d", table.Rows, n)
|
||||
}
|
||||
if table.Headers["ORIGIN"] != "tensor test" {
|
||||
t.Fatalf("header ORIGIN = %q", table.Headers["ORIGIN"])
|
||||
}
|
||||
wantNames := []string{"ID", "MAG", "TEMP"}
|
||||
for i, want := range wantNames {
|
||||
if table.Names[i] != want {
|
||||
t.Fatalf("column %d named %q, want %q", i, table.Names[i], want)
|
||||
}
|
||||
}
|
||||
if table.Units[1] != "mag" {
|
||||
t.Fatalf("MAG unit = %q, want mag", table.Units[1])
|
||||
}
|
||||
for i := range n {
|
||||
if table.Columns[0].RawInts()[i] != int64(1000+i) {
|
||||
t.Fatalf("ID[%d] = %d", i, table.Columns[0].RawInts()[i])
|
||||
}
|
||||
if math.Abs(table.Columns[1].FloatAt(i)-cols[1].Data.FloatAt(i)) > 1e-12 {
|
||||
t.Fatalf("MAG[%d] = %.14g", i, table.Columns[1].FloatAt(i))
|
||||
}
|
||||
if math.Abs(float64(table.Columns[2].RawFloat32s()[i]-cols[2].Data.RawFloat32s()[i])) > 1e-4 {
|
||||
t.Fatalf("TEMP[%d] = %.6g", i, table.Columns[2].RawFloat32s()[i])
|
||||
}
|
||||
}
|
||||
// The primary image still reads through the image loader.
|
||||
img, headers, err := LoadFITS(path)
|
||||
if err != nil {
|
||||
t.Fatalf("LoadFITS on a table file: %v", err)
|
||||
}
|
||||
if img.Len() != 0 {
|
||||
t.Fatalf("primary image has %d elements, want an empty zero-axis HDU", img.Len())
|
||||
}
|
||||
if headers["EXTEND"] != "T" {
|
||||
t.Fatalf("EXTEND = %q, want T", headers["EXTEND"])
|
||||
}
|
||||
}
|
||||
|
||||
// TestSaveLoadFITSTableBinaryStringColumnNotFirst round-trips a
|
||||
// binary table whose character column sits between two numeric ones.
|
||||
// A character encoder that writes at the start of the row instead of
|
||||
// its own column offset corrupts every column beside it, and the
|
||||
// damage is silent: the file still loads.
|
||||
func TestSaveLoadFITSTableBinaryStringColumnNotFirst(t *testing.T) {
|
||||
const n = 3
|
||||
mags := []float64{1.25, -0.5, 3.75}
|
||||
names := []string{"alf", "bet", "gam"}
|
||||
ids := []int64{7, 8, 9}
|
||||
cols := []FITSTableColumn{
|
||||
{Name: "MAG", Unit: "mag", Form: "D", Data: mustFloats(t, mags, n)},
|
||||
{Name: "STAR", Form: "8A", Text: names},
|
||||
{Name: "ID", Form: "K", Data: mustInts(t, ids, n)},
|
||||
}
|
||||
path := filepath.Join(t.TempDir(), "stars.fits")
|
||||
if err := SaveFITSTable(path, false, cols, nil); err != nil {
|
||||
t.Fatalf("SaveFITSTable: %v", err)
|
||||
}
|
||||
table, err := LoadFITSTable(path)
|
||||
if err != nil {
|
||||
t.Fatalf("LoadFITSTable: %v", err)
|
||||
}
|
||||
for i := range n {
|
||||
if got := table.Columns[0].FloatAt(i); got != mags[i] {
|
||||
t.Fatalf("MAG[%d] = %v, want %v", i, got, mags[i])
|
||||
}
|
||||
if got := table.Text[1][i]; got != names[i] {
|
||||
t.Fatalf("STAR[%d] = %q, want %q", i, got, names[i])
|
||||
}
|
||||
if got := table.Columns[2].RawInts()[i]; got != ids[i] {
|
||||
t.Fatalf("ID[%d] = %d, want %d", i, got, ids[i])
|
||||
}
|
||||
}
|
||||
// The row layout itself: 8 bytes of float64, 8 of text, 8 of int.
|
||||
raw, err := os.ReadFile(path)
|
||||
if err != nil {
|
||||
t.Fatalf("read back: %v", err)
|
||||
}
|
||||
start := bytes.Index(raw, []byte("XTENSION"))
|
||||
if start < 0 {
|
||||
t.Fatal("no XTENSION card in the file")
|
||||
}
|
||||
// The data starts on the 2880-byte boundary after the extension
|
||||
// header, whose card count the loader has already validated.
|
||||
data := raw[(start/2880+1)*2880:]
|
||||
if got := math.Float64frombits(binary.BigEndian.Uint64(data[0:])); got != mags[0] {
|
||||
t.Fatalf("first row's first field = %v, want %v", got, mags[0])
|
||||
}
|
||||
if got := string(bytes.TrimRight(data[8:16], " ")); got != names[0] {
|
||||
t.Fatalf("first row's text field = %q, want %q", got, names[0])
|
||||
}
|
||||
if got := int64(binary.BigEndian.Uint64(data[16:24])); got != ids[0] {
|
||||
t.Fatalf("first row's int field = %d, want %d", got, ids[0])
|
||||
}
|
||||
}
|
||||
|
||||
// TestSaveLoadFITSTableASCII round-trips an ASCII table with integer,
|
||||
// float and string columns.
|
||||
func TestSaveLoadFITSTableASCII(t *testing.T) {
|
||||
const n = 4
|
||||
names := []string{"ALF Cen", "Betel", "Rigel", "Deneb"}
|
||||
mags := make([]float64, n)
|
||||
for i := range n {
|
||||
mags[i] = -1.5 + 1.3*float64(i)
|
||||
}
|
||||
cols := []FITSTableColumn{
|
||||
{Name: "STAR", Form: "10A", Text: names},
|
||||
{Name: "MAG", Unit: "mag", Form: "D20.14", Data: mustFloats(t, mags, n)},
|
||||
}
|
||||
path := filepath.Join(t.TempDir(), "stars.fits")
|
||||
if err := SaveFITSTable(path, true, cols, nil); err != nil {
|
||||
t.Fatalf("SaveFITSTable: %v", err)
|
||||
}
|
||||
table, err := LoadFITSTable(path)
|
||||
if err != nil {
|
||||
t.Fatalf("LoadFITSTable: %v", err)
|
||||
}
|
||||
if table.Kind != "TABLE" {
|
||||
t.Fatalf("kind %q, want TABLE", table.Kind)
|
||||
}
|
||||
for i := range n {
|
||||
if table.Text[0][i] != names[i] {
|
||||
t.Fatalf("STAR[%d] = %q, want %q", i, table.Text[0][i], names[i])
|
||||
}
|
||||
if math.Abs(table.Columns[1].FloatAt(i)-mags[i]) > 1e-9 {
|
||||
t.Fatalf("MAG[%d] = %.14g, want %.14g", i, table.Columns[1].FloatAt(i), mags[i])
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// TestLoadFITSTableSkipsImage writes an image first and a table
|
||||
// second by hand-concatenating the two files' HDUs: the loader must
|
||||
// skip the image and land on the table.
|
||||
func TestLoadFITSTableSkipsImage(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
img := mustFloats(t, []float64{1, 2, 3, 4}, 2, 2)
|
||||
imagePath := filepath.Join(dir, "image.fits")
|
||||
if err := SaveFITS(imagePath, img, nil); err != nil {
|
||||
t.Fatalf("SaveFITS: %v", err)
|
||||
}
|
||||
cols := []FITSTableColumn{
|
||||
{Name: "X", Form: "D", Data: mustFloats(t, []float64{1.5, 2.5}, 2)},
|
||||
}
|
||||
tablePath := filepath.Join(dir, "table.fits")
|
||||
if err := SaveFITSTable(tablePath, false, cols, nil); err != nil {
|
||||
t.Fatalf("SaveFITSTable: %v", err)
|
||||
}
|
||||
combined := filepath.Join(dir, "combined.fits")
|
||||
imageData, err := osReadFile(imagePath)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
tableData, err := osReadFile(tablePath)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
// The image file's primary already declares EXTEND; append the
|
||||
// table extension with its own primary stripped (the extension
|
||||
// starts at its XTENSION card, 2880 bytes in).
|
||||
if err := osWriteFile(combined, imageData, tableData[2880:]); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
table, err := LoadFITSTable(combined)
|
||||
if err != nil {
|
||||
t.Fatalf("LoadFITSTable: %v", err)
|
||||
}
|
||||
if table.Columns[0].FloatAt(1) != 2.5 {
|
||||
t.Fatalf("X[1] = %.4g, want 2.5", table.Columns[0].FloatAt(1))
|
||||
}
|
||||
}
|
||||
|
||||
// TestSaveFITSTableErrors pins the validation contract.
|
||||
func TestSaveFITSTableErrors(t *testing.T) {
|
||||
if err := SaveFITSTable(filepath.Join(t.TempDir(), "x.fits"), false, nil, nil); err == nil {
|
||||
t.Fatal("expected an error for an empty column list")
|
||||
}
|
||||
noData := []FITSTableColumn{{Name: "X", Form: "D"}}
|
||||
if err := SaveFITSTable(filepath.Join(t.TempDir(), "x.fits"), false, noData, nil); err == nil {
|
||||
t.Fatal("expected an error for a numeric column without data")
|
||||
}
|
||||
noText := []FITSTableColumn{{Name: "S", Form: "8A"}}
|
||||
if err := SaveFITSTable(filepath.Join(t.TempDir(), "x.fits"), false, noText, nil); err == nil {
|
||||
t.Fatal("expected an error for a character column without text")
|
||||
}
|
||||
badForm := []FITSTableColumn{{Name: "X", Form: "Q", Data: mustFloats(t, []float64{1}, 1)}}
|
||||
if err := SaveFITSTable(filepath.Join(t.TempDir(), "x.fits"), false, badForm, nil); err == nil {
|
||||
t.Fatal("expected an error for an unknown form")
|
||||
}
|
||||
ragged := []FITSTableColumn{
|
||||
{Name: "X", Form: "D", Data: mustFloats(t, []float64{1, 2}, 2)},
|
||||
{Name: "Y", Form: "D", Data: mustFloats(t, []float64{1}, 1)},
|
||||
}
|
||||
if err := SaveFITSTable(filepath.Join(t.TempDir(), "x.fits"), false, ragged, nil); err == nil {
|
||||
t.Fatal("expected an error for mismatched row counts")
|
||||
}
|
||||
longName := []FITSTableColumn{{Name: "S", Form: "4A", Text: []string{"too long"}}}
|
||||
if err := SaveFITSTable(filepath.Join(t.TempDir(), "x.fits"), false, longName, nil); err == nil {
|
||||
t.Fatal("expected an error for text wider than the form")
|
||||
}
|
||||
}
|
||||
|
||||
// TestLoadFITSTableHostileRows pins the NAXIS2 guards: a hostile
|
||||
// header declaring a colossal row count must report truncation (not
|
||||
// overflow or an out-of-memory allocation), and a negative one must
|
||||
// be rejected instead of answering an empty table.
|
||||
func TestLoadFITSTableHostileRows(t *testing.T) {
|
||||
cols, n := starCatalogue(t)
|
||||
path := filepath.Join(t.TempDir(), "catalogue.fits")
|
||||
if err := SaveFITSTable(path, false, cols, nil); err != nil {
|
||||
t.Fatalf("SaveFITSTable: %v", err)
|
||||
}
|
||||
raw, err := osReadFile(path)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
hostile := filepath.Join(t.TempDir(), "hostile.fits")
|
||||
huge := osWriteFile(hostile, bytes.Replace(raw, []byte(fitsIntCard("NAXIS2", n)), []byte(fitsIntCard("NAXIS2", math.MaxInt64)), 1))
|
||||
if huge != nil {
|
||||
t.Fatalf("os.WriteFile: %v", huge)
|
||||
}
|
||||
if _, err := LoadFITSTable(hostile); err == nil {
|
||||
t.Fatal("expected an error for a colossal NAXIS2")
|
||||
}
|
||||
negative := filepath.Join(t.TempDir(), "negative.fits")
|
||||
if err := osWriteFile(negative, bytes.Replace(raw, []byte(fitsIntCard("NAXIS2", n)), []byte(fitsIntCard("NAXIS2", -5)), 1)); err != nil {
|
||||
t.Fatalf("os.WriteFile: %v", err)
|
||||
}
|
||||
if _, err := LoadFITSTable(negative); err == nil {
|
||||
t.Fatal("expected an error for a negative NAXIS2")
|
||||
}
|
||||
}
|
||||
|
||||
// TestSaveFITSTableASCIIWidthErrors pins the ASCII field contract: a
|
||||
// value whose rendering exceeds the form's declared width is an
|
||||
// error, never silently truncated digits or an overflowing field.
|
||||
func TestSaveFITSTableASCIIWidthErrors(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
wideInt := []FITSTableColumn{{Name: "N", Form: "I5", Data: mustInts(t, []int64{12345678}, 1)}}
|
||||
if err := SaveFITSTable(filepath.Join(dir, "i.fits"), true, wideInt, nil); err == nil {
|
||||
t.Error("expected an error for an integer that does not fit I5")
|
||||
}
|
||||
wideFloat := []FITSTableColumn{{Name: "F", Form: "F10.4", Data: mustFloats(t, []float64{1e300}, 1)}}
|
||||
if err := SaveFITSTable(filepath.Join(dir, "f.fits"), true, wideFloat, nil); err == nil {
|
||||
t.Error("expected an error for a float that does not fit F10.4")
|
||||
}
|
||||
wideText := []FITSTableColumn{{Name: "S", Form: "3A", Text: []string{"abcd"}}}
|
||||
if err := SaveFITSTable(filepath.Join(dir, "a.fits"), true, wideText, nil); err == nil {
|
||||
t.Error("expected an error for text wider than the form")
|
||||
}
|
||||
// Fitting values keep working.
|
||||
okCols := []FITSTableColumn{
|
||||
{Name: "N", Form: "I5", Data: mustInts(t, []int64{42}, 1)},
|
||||
{Name: "F", Form: "F10.4", Data: mustFloats(t, []float64{3.14159}, 1)},
|
||||
{Name: "E", Form: "E13.4", Data: mustFloats(t, []float64{-1.5e300}, 1)},
|
||||
}
|
||||
path := filepath.Join(dir, "ok.fits")
|
||||
if err := SaveFITSTable(path, true, okCols, nil); err != nil {
|
||||
t.Fatalf("SaveFITSTable with fitting values: %v", err)
|
||||
}
|
||||
table, err := LoadFITSTable(path)
|
||||
if err != nil {
|
||||
t.Fatalf("LoadFITSTable: %v", err)
|
||||
}
|
||||
if table.Columns[0].RawInts()[0] != 42 {
|
||||
t.Errorf("N[0] = %d, want 42", table.Columns[0].RawInts()[0])
|
||||
}
|
||||
if math.Abs(table.Columns[1].FloatAt(0)-3.14159) > 1e-4 {
|
||||
t.Errorf("F[0] = %.6g, want 3.14159", table.Columns[1].FloatAt(0))
|
||||
}
|
||||
if got, want := table.Columns[2].FloatAt(0), -1.5e300; math.Abs(got-want) > 1e-6*math.Abs(want) {
|
||||
t.Errorf("E[0] = %.6g, want %.6g", got, want)
|
||||
}
|
||||
}
|
||||
|
||||
// TestLoadFITSTableSignedBinaryForms pins the four numeric decodes the
|
||||
// package's own writer cannot emit: J is a signed 32-bit big-endian
|
||||
// integer, I a signed 16-bit one, B an unsigned byte and L a logical
|
||||
// whose false value is the zero the payload was allocated with. The
|
||||
// table is laid out here, so every value that distinguishes the forms
|
||||
// is present: a negative integer in each of J and I, a byte above the
|
||||
// signed range and a false logical beside a true one.
|
||||
func TestLoadFITSTableSignedBinaryForms(t *testing.T) {
|
||||
hdr := cardBlock(
|
||||
card("XTENSION= 'BINTABLE'"),
|
||||
card("BITPIX = 8"),
|
||||
card("NAXIS = 2"),
|
||||
card("NAXIS1 = 8"),
|
||||
card("NAXIS2 = 3"),
|
||||
card("TFIELDS = 4"),
|
||||
card("TTYPE1 = 'JCOL '"),
|
||||
card("TFORM1 = 'J '"),
|
||||
card("TTYPE2 = 'ICOL '"),
|
||||
card("TFORM2 = 'I '"),
|
||||
card("TTYPE3 = 'BCOL '"),
|
||||
card("TFORM3 = 'B '"),
|
||||
card("TTYPE4 = 'LCOL '"),
|
||||
card("TFORM4 = 'L '"),
|
||||
card("END"),
|
||||
)
|
||||
rows := []struct {
|
||||
j int32
|
||||
i int16
|
||||
b byte
|
||||
l byte
|
||||
}{
|
||||
{-123456, -7, 200, 'T'},
|
||||
{123456, 30000, 255, 'F'},
|
||||
{-1, -32768, 128, 'T'},
|
||||
}
|
||||
var body []byte
|
||||
for _, r := range rows {
|
||||
body = binary.BigEndian.AppendUint32(body, uint32(r.j))
|
||||
body = binary.BigEndian.AppendUint16(body, uint16(r.i))
|
||||
body = append(body, r.b, r.l)
|
||||
}
|
||||
path := writeHostile(t, "forms.fits", append(hdr, body...))
|
||||
table, err := LoadFITSTable(path)
|
||||
if err != nil {
|
||||
t.Fatalf("LoadFITSTable: %v", err)
|
||||
}
|
||||
if len(table.Columns) != 4 {
|
||||
t.Fatalf("the table carries %d columns, want 4", len(table.Columns))
|
||||
}
|
||||
for row, want := range rows {
|
||||
if got := table.Columns[0].RawInts()[row]; got != int64(want.j) {
|
||||
t.Fatalf("row %d: J = %d, want %d (a signed 32-bit big-endian integer)", row, got, want.j)
|
||||
}
|
||||
if got := table.Columns[1].RawInts()[row]; got != int64(want.i) {
|
||||
t.Fatalf("row %d: I = %d, want %d (a signed 16-bit big-endian integer)", row, got, want.i)
|
||||
}
|
||||
if got := table.Columns[2].RawInts()[row]; got != int64(want.b) {
|
||||
t.Fatalf("row %d: B = %d, want %d (an unsigned byte)", row, got, want.b)
|
||||
}
|
||||
wantL := int64(0)
|
||||
if want.l == 'T' {
|
||||
wantL = 1
|
||||
}
|
||||
if got := table.Columns[3].RawInts()[row]; got != wantL {
|
||||
t.Fatalf("row %d: L = %d, want %d (%q in the file)", row, got, wantL, want.l)
|
||||
}
|
||||
}
|
||||
}
|
||||
+474
@@ -0,0 +1,474 @@
|
||||
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
|
||||
// SPDX-License-Identifier: MIT
|
||||
|
||||
package io
|
||||
|
||||
import (
|
||||
"math"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"slices"
|
||||
"testing"
|
||||
|
||||
"sourcedock.dev/petrbalvin/tensor/internal/core"
|
||||
)
|
||||
|
||||
// Coverage-guided fuzzing over the four binary readers. Every input is
|
||||
// written to one shared temp file per worker process and handed to the
|
||||
// reader; a panic, a hang or an allocation blow-up fails the input, and
|
||||
// so does a returned result whose shape disagrees with its payload,
|
||||
// because a silent mismatch is the worst outcome a parser can produce.
|
||||
// Seeds are the real fixtures plus files built by the own writers, so
|
||||
// the mutator starts from genuinely valid bytes rather than from magic
|
||||
// constants. Run one target at a time, for example:
|
||||
//
|
||||
// go test -run '^$' -fuzz FuzzLoadHDF5 -fuzztime 2m ./io/
|
||||
//
|
||||
// Inputs that expose a defect land in testdata/fuzz/<FuzzName>/ and
|
||||
// stay there as regression seeds.
|
||||
|
||||
// fuzzWrite seeds the corpus with one file built by a writer.
|
||||
func fuzzWrite(f *testing.F, name string, build func(path string) error) {
|
||||
path := filepath.Join(f.TempDir(), name)
|
||||
if err := build(path); err != nil {
|
||||
f.Fatalf("build seed %s: %v", name, err)
|
||||
}
|
||||
data, err := os.ReadFile(path)
|
||||
if err != nil {
|
||||
f.Fatalf("read seed %s: %v", name, err)
|
||||
}
|
||||
f.Add(data)
|
||||
}
|
||||
|
||||
// fuzzFloats builds a float64 vector of n recognisably distinct values.
|
||||
func fuzzFloats(f *testing.F, n int) *core.Array {
|
||||
f.Helper()
|
||||
vals := make([]float64, n)
|
||||
for i := range vals {
|
||||
vals[i] = float64(i+1) + 0.5
|
||||
}
|
||||
a, err := core.FromFloats(vals, n)
|
||||
if err != nil {
|
||||
f.Fatalf("FromFloats(%d): %v", n, err)
|
||||
}
|
||||
return a
|
||||
}
|
||||
|
||||
// fuzzInts builds an int64 vector of n recognisably distinct values.
|
||||
func fuzzInts(f *testing.F, n int) *core.Array {
|
||||
f.Helper()
|
||||
vals := make([]int64, n)
|
||||
for i := range vals {
|
||||
vals[i] = int64(i) - 2
|
||||
}
|
||||
a, err := core.FromInts(vals, n)
|
||||
if err != nil {
|
||||
f.Fatalf("FromInts(%d): %v", n, err)
|
||||
}
|
||||
return a
|
||||
}
|
||||
|
||||
// fuzzShapeProduct multiplies the shape with the overflow guard the
|
||||
// parsers themselves are held to: an attacker-controlled shape must
|
||||
// never wrap around into a plausible small length.
|
||||
func fuzzShapeProduct(t *testing.T, shape []int) int {
|
||||
n := 1
|
||||
for _, d := range shape {
|
||||
if d < 0 || d > (1<<31-1)/n {
|
||||
t.Fatalf("hostile shape %v accepted", shape)
|
||||
}
|
||||
n *= d
|
||||
}
|
||||
return n
|
||||
}
|
||||
|
||||
func FuzzLoadHDF5(f *testing.F) {
|
||||
fuzzWrite(f, "fixture.h5", func(p string) error {
|
||||
data, err := os.ReadFile("testdata/h5/fixture.h5")
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return os.WriteFile(p, data, 0o644)
|
||||
})
|
||||
fuzzWrite(f, "fletcher.h5", func(p string) error {
|
||||
data, err := os.ReadFile("testdata/h5/fletcher.h5")
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return os.WriteFile(p, data, 0o644)
|
||||
})
|
||||
fuzzWrite(f, "latest.h5", func(p string) error {
|
||||
data, err := os.ReadFile("testdata/h5/latest.h5")
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return os.WriteFile(p, data, 0o644)
|
||||
})
|
||||
path := filepath.Join(f.TempDir(), "fuzz.h5")
|
||||
f.Fuzz(func(t *testing.T, data []byte) {
|
||||
if err := os.WriteFile(path, data, 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
sets, err := LoadHDF5(path)
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
for _, s := range sets {
|
||||
if n := fuzzShapeProduct(t, s.Shape); n != s.Values.Len() {
|
||||
t.Fatalf("dataset %s: shape product %d != payload %d", s.Path, n, s.Values.Len())
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func FuzzLoadNetCDF(f *testing.F) {
|
||||
fuzzWrite(f, "classic.nc", func(p string) error {
|
||||
dims := []NetCDFDim{{Name: "lat", Length: 3}, {Name: "lon", Length: 4}}
|
||||
vars := []NetCDFVar{
|
||||
{Name: "temp", Dims: []string{"lat", "lon"}, Values: fuzzFloats(f, 12)},
|
||||
{Name: "mask", Dims: []string{"lon"}, Values: fuzzInts(f, 4)},
|
||||
}
|
||||
return SaveNetCDF(p, dims, vars, map[string]string{"title": "seed"})
|
||||
})
|
||||
fuzzWrite(f, "record.nc", func(p string) error {
|
||||
dims := []NetCDFDim{{Name: "t", Length: 0}, {Name: "x", Length: 2}}
|
||||
vars := []NetCDFVar{
|
||||
{Name: "series", Dims: []string{"t", "x"}, Values: fuzzFloats(f, 6)},
|
||||
}
|
||||
return SaveNetCDF(p, dims, vars, nil)
|
||||
})
|
||||
path := filepath.Join(f.TempDir(), "fuzz.nc")
|
||||
f.Fuzz(func(t *testing.T, data []byte) {
|
||||
if err := os.WriteFile(path, data, 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
dims, vars, _, err := LoadNetCDF(path)
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
lengths := make(map[string]int, len(dims))
|
||||
for _, d := range dims {
|
||||
lengths[d.Name] = d.Length
|
||||
}
|
||||
for _, v := range vars {
|
||||
// A rank-0 (scalar) variable comes back as a one-element
|
||||
// vector with no dimension names, the documented shape.
|
||||
if len(v.Values.Shape()) != len(v.Dims) && !(len(v.Dims) == 0 && v.Values.Len() == 1) {
|
||||
t.Fatalf("variable %s: rank %d != %d named dimensions", v.Name, len(v.Values.Shape()), len(v.Dims))
|
||||
}
|
||||
record := false
|
||||
n := 1
|
||||
for _, name := range v.Dims {
|
||||
length, ok := lengths[name]
|
||||
if !ok {
|
||||
t.Fatalf("variable %s: unknown dimension %s", v.Name, name)
|
||||
}
|
||||
if length == 0 {
|
||||
record = true
|
||||
continue
|
||||
}
|
||||
n *= length
|
||||
}
|
||||
if !record && n != v.Values.Len() {
|
||||
t.Fatalf("variable %s: dimension product %d != payload %d", v.Name, n, v.Values.Len())
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func FuzzLoadFITS(f *testing.F) {
|
||||
fuzzWrite(f, "image.fits", func(p string) error {
|
||||
a, err := core.FromFloats([]float64{1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15, 16}, 4, 4)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return SaveFITS(p, a, map[string]string{"OBJECT": "seed", "TELESCOP": "TENSOR"})
|
||||
})
|
||||
path := filepath.Join(f.TempDir(), "fuzz.fits")
|
||||
f.Fuzz(func(t *testing.T, data []byte) {
|
||||
if err := os.WriteFile(path, data, 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
a, _, err := LoadFITS(path)
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
if n := fuzzShapeProduct(t, a.Shape()); n != a.Len() {
|
||||
t.Fatalf("image: shape product %d != payload %d", n, a.Len())
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func FuzzLoadFITSTable(f *testing.F) {
|
||||
fuzzWrite(f, "binary.fits", func(p string) error {
|
||||
cols := []FITSTableColumn{
|
||||
{Name: "flux", Form: "D", Data: fuzzFloats(f, 5)},
|
||||
{Name: "id", Form: "K", Data: fuzzInts(f, 5)},
|
||||
{Name: "name", Form: "8A", Text: []string{"alpha", "beta", "gamma", "delta", "epsilon"}},
|
||||
}
|
||||
return SaveFITSTable(p, false, cols, nil)
|
||||
})
|
||||
fuzzWrite(f, "ascii.fits", func(p string) error {
|
||||
cols := []FITSTableColumn{
|
||||
{Name: "flux", Form: "E12.5", Data: fuzzFloats(f, 3)},
|
||||
{Name: "name", Form: "6A", Text: []string{"one", "two", "three"}},
|
||||
}
|
||||
return SaveFITSTable(p, true, cols, nil)
|
||||
})
|
||||
path := filepath.Join(f.TempDir(), "fuzztable.fits")
|
||||
f.Fuzz(func(t *testing.T, data []byte) {
|
||||
if err := os.WriteFile(path, data, 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
table, err := LoadFITSTable(path)
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
if table.Rows < 0 {
|
||||
t.Fatalf("negative row count %d", table.Rows)
|
||||
}
|
||||
n := len(table.Names)
|
||||
if len(table.Columns) != n || len(table.Text) != n {
|
||||
t.Fatalf("column lists disagree: %d names, %d value columns, %d text columns", n, len(table.Columns), len(table.Text))
|
||||
}
|
||||
for i := range table.Names {
|
||||
switch {
|
||||
case table.Columns[i] != nil && table.Text[i] != nil:
|
||||
t.Fatalf("column %d carries both values and text", i)
|
||||
case table.Columns[i] != nil:
|
||||
if table.Columns[i].Len() != table.Rows {
|
||||
t.Fatalf("column %d: length %d != %d rows", i, table.Columns[i].Len(), table.Rows)
|
||||
}
|
||||
case table.Text[i] != nil:
|
||||
if len(table.Text[i]) != table.Rows {
|
||||
t.Fatalf("text column %d: length %d != %d rows", i, len(table.Text[i]), table.Rows)
|
||||
}
|
||||
default:
|
||||
t.Fatalf("column %d carries neither values nor text", i)
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
// FuzzHDF5WriteRead fuzzes the writer's round-trip contract: whatever
|
||||
// parameters and values the mutator picks, SaveHDF5 must produce a
|
||||
// file the reader decodes back to the very same shapes, bit patterns
|
||||
// and attributes. A write that fails, a read that fails or a value
|
||||
// that moves is a bug, not a skipped input: unlike the readers there
|
||||
// is no untrusted bytes here, the writer owns every byte it emits.
|
||||
func FuzzHDF5WriteRead(f *testing.F) {
|
||||
for _, seed := range [][]byte{
|
||||
{0, 0, 7, 3, 0, 0, 0, 0},
|
||||
{0, 0, 7, 3, 0, 0, 1, 0},
|
||||
{1, 1, 19, 5, 1, 1, 0, 1},
|
||||
{2, 0, 33, 1, 0, 1, 1, 1},
|
||||
{0, 1, 5, 7, 1, 0, 0, 2},
|
||||
{2, 1, 11, 4, 1, 1, 1, 1},
|
||||
{3, 0, 5, 4, 0, 0, 0, 0},
|
||||
{4, 1, 9, 3, 0, 1, 0, 1},
|
||||
{5, 0, 6, 2, 0, 0, 1, 0},
|
||||
{6, 1, 12, 5, 0, 0, 0, 1},
|
||||
{7, 0, 8, 3, 1, 1, 0, 0},
|
||||
{8, 1, 15, 4, 0, 0, 0, 1},
|
||||
{9, 0, 10, 2, 1, 0, 0, 0},
|
||||
} {
|
||||
f.Add(seed)
|
||||
}
|
||||
f.Fuzz(func(t *testing.T, data []byte) {
|
||||
if len(data) < 8 {
|
||||
t.Skip()
|
||||
}
|
||||
dtype := int(data[0] % 10)
|
||||
rank := 1 + int(data[1]%2)
|
||||
d0 := 1 + int(data[2])%31
|
||||
d1 := 1 + int(data[3])%9
|
||||
gzip := 0
|
||||
if data[4]%2 == 1 {
|
||||
gzip = 6
|
||||
}
|
||||
shuffle := data[5]%2 == 1
|
||||
// A filtered dataset in a latest-version file is refused by
|
||||
// design; the fuzz contract covers the accepted combinations.
|
||||
latest := data[6]%2 == 1 && gzip == 0 && !shuffle
|
||||
chunk := 0
|
||||
if data[7]%2 == 1 {
|
||||
chunk = 32
|
||||
}
|
||||
n := d0
|
||||
if rank == 2 {
|
||||
n *= d1
|
||||
}
|
||||
shape := []int{d0}
|
||||
if rank == 2 {
|
||||
shape = append(shape, d1)
|
||||
}
|
||||
var values *core.Array
|
||||
var err error
|
||||
switch dtype {
|
||||
case 0:
|
||||
v := make([]float64, n)
|
||||
for i := range v {
|
||||
v[i] = float64(data[(8+i)%len(data)]) - 128 + float64(i%5)*0.25
|
||||
}
|
||||
values, err = core.FromFloats(v, shape...)
|
||||
case 1:
|
||||
v := make([]float32, n)
|
||||
for i := range v {
|
||||
v[i] = float32(data[(8+i)%len(data)]) - 128 + float32(i%5)*0.25
|
||||
}
|
||||
values, err = core.FromFloat32s(v, shape...)
|
||||
case 2:
|
||||
v := make([]int64, n)
|
||||
for i := range v {
|
||||
v[i] = int64(data[(8+i)%len(data)])*1000 - 128000 + int64(i)
|
||||
}
|
||||
values, err = core.FromInts(v, shape...)
|
||||
case 3:
|
||||
v := make([]bool, n)
|
||||
for i := range v {
|
||||
v[i] = data[(8+i)%len(data)]%2 == 1
|
||||
}
|
||||
values, err = core.FromBools(v, shape...)
|
||||
case 4:
|
||||
v := make([]int8, n)
|
||||
for i := range v {
|
||||
v[i] = int8(int(data[(8+i)%len(data)]) - 128)
|
||||
}
|
||||
values, err = core.FromInt8s(v, shape...)
|
||||
case 5:
|
||||
v := make([]uint8, n)
|
||||
for i := range v {
|
||||
v[i] = data[(8+i)%len(data)]
|
||||
}
|
||||
values, err = core.FromUint8s(v, shape...)
|
||||
case 6:
|
||||
v := make([]int16, n)
|
||||
for i := range v {
|
||||
v[i] = int16(int(data[(8+i)%len(data)])*257 - 32768)
|
||||
}
|
||||
values, err = core.FromInt16s(v, shape...)
|
||||
case 7:
|
||||
v := make([]uint16, n)
|
||||
for i := range v {
|
||||
v[i] = uint16(int(data[(8+i)%len(data)]) * 257)
|
||||
}
|
||||
values, err = core.FromUint16s(v, shape...)
|
||||
case 8:
|
||||
v := make([]int32, n)
|
||||
for i := range v {
|
||||
v[i] = int32(int64(data[(8+i)%len(data)])*1000000 - 128000000 + int64(i))
|
||||
}
|
||||
values, err = core.FromInt32s(v, shape...)
|
||||
default:
|
||||
v := make([]uint32, n)
|
||||
for i := range v {
|
||||
v[i] = uint32(uint64(data[(8+i)%len(data)])*1000000 + uint64(i))
|
||||
}
|
||||
values, err = core.FromUint32s(v, shape...)
|
||||
}
|
||||
if err != nil {
|
||||
t.Fatalf("build input: %v", err)
|
||||
}
|
||||
sets := []HDF5Dataset{{Path: "/d", Shape: shape, Values: values, Attrs: map[string]string{"k": "1"}}}
|
||||
attrs := map[string]map[string]string{"/": {"title": "fuzz"}}
|
||||
path := filepath.Join(t.TempDir(), "fuzz.h5")
|
||||
if err := SaveHDF5(path, sets, attrs, HDF5WriteOptions{Gzip: gzip, Shuffle: shuffle, Latest: latest, ChunkBytes: chunk}); err != nil {
|
||||
t.Fatalf("SaveHDF5 (%v): %v", shape, err)
|
||||
}
|
||||
back, err := LoadHDF5(path)
|
||||
if err != nil {
|
||||
t.Fatalf("LoadHDF5 round trip (%v): %v", shape, err)
|
||||
}
|
||||
if len(back) != 1 {
|
||||
t.Fatalf("round trip returned %d datasets, want 1", len(back))
|
||||
}
|
||||
d := back[0]
|
||||
if d.Path != "/d" {
|
||||
t.Fatalf("path = %q, want /d", d.Path)
|
||||
}
|
||||
if !slices.Equal(d.Shape, shape) {
|
||||
t.Fatalf("shape = %v, want %v", d.Shape, shape)
|
||||
}
|
||||
if d.Values.Dtype() != values.Dtype() {
|
||||
t.Fatalf("dtype = %s, want %s", d.Values.Dtype(), values.Dtype())
|
||||
}
|
||||
switch dtype {
|
||||
case 0:
|
||||
got, want := d.Values.RawFloats()[:n], values.RawFloats()[:n]
|
||||
for i := range want {
|
||||
if math.Float64bits(got[i]) != math.Float64bits(want[i]) {
|
||||
t.Fatalf("[%d] = %v, want %v (bit-exact)", i, got[i], want[i])
|
||||
}
|
||||
}
|
||||
case 1:
|
||||
got, want := d.Values.RawFloat32s()[:n], values.RawFloat32s()[:n]
|
||||
for i := range want {
|
||||
if math.Float32bits(got[i]) != math.Float32bits(want[i]) {
|
||||
t.Fatalf("[%d] = %v, want %v (bit-exact)", i, got[i], want[i])
|
||||
}
|
||||
}
|
||||
case 2:
|
||||
got, want := d.Values.RawInts()[:n], values.RawInts()[:n]
|
||||
for i := range want {
|
||||
if got[i] != want[i] {
|
||||
t.Fatalf("[%d] = %d, want %d", i, got[i], want[i])
|
||||
}
|
||||
}
|
||||
case 3:
|
||||
got, want := d.Values.RawBools()[:n], values.RawBools()[:n]
|
||||
for i := range want {
|
||||
if got[i] != want[i] {
|
||||
t.Fatalf("[%d] = %v, want %v", i, got[i], want[i])
|
||||
}
|
||||
}
|
||||
case 4:
|
||||
got, want := d.Values.RawInt8s()[:n], values.RawInt8s()[:n]
|
||||
for i := range want {
|
||||
if got[i] != want[i] {
|
||||
t.Fatalf("[%d] = %d, want %d", i, got[i], want[i])
|
||||
}
|
||||
}
|
||||
case 5:
|
||||
got, want := d.Values.RawUint8s()[:n], values.RawUint8s()[:n]
|
||||
for i := range want {
|
||||
if got[i] != want[i] {
|
||||
t.Fatalf("[%d] = %d, want %d", i, got[i], want[i])
|
||||
}
|
||||
}
|
||||
case 6:
|
||||
got, want := d.Values.RawInt16s()[:n], values.RawInt16s()[:n]
|
||||
for i := range want {
|
||||
if got[i] != want[i] {
|
||||
t.Fatalf("[%d] = %d, want %d", i, got[i], want[i])
|
||||
}
|
||||
}
|
||||
case 7:
|
||||
got, want := d.Values.RawUint16s()[:n], values.RawUint16s()[:n]
|
||||
for i := range want {
|
||||
if got[i] != want[i] {
|
||||
t.Fatalf("[%d] = %d, want %d", i, got[i], want[i])
|
||||
}
|
||||
}
|
||||
case 8:
|
||||
got, want := d.Values.RawInt32s()[:n], values.RawInt32s()[:n]
|
||||
for i := range want {
|
||||
if got[i] != want[i] {
|
||||
t.Fatalf("[%d] = %d, want %d", i, got[i], want[i])
|
||||
}
|
||||
}
|
||||
default:
|
||||
got, want := d.Values.RawUint32s()[:n], values.RawUint32s()[:n]
|
||||
for i := range want {
|
||||
if got[i] != want[i] {
|
||||
t.Fatalf("[%d] = %d, want %d", i, got[i], want[i])
|
||||
}
|
||||
}
|
||||
}
|
||||
if d.Attrs["k"] != "1" {
|
||||
t.Fatalf("dataset attr k = %q, want 1", d.Attrs["k"])
|
||||
}
|
||||
if d.Attrs["title"] != "fuzz" {
|
||||
t.Fatalf("merged root attr title = %q, want fuzz", d.Attrs["title"])
|
||||
}
|
||||
})
|
||||
}
|
||||
+2303
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,90 @@
|
||||
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
|
||||
// SPDX-License-Identifier: MIT
|
||||
|
||||
package io
|
||||
|
||||
import (
|
||||
"path/filepath"
|
||||
"slices"
|
||||
"testing"
|
||||
|
||||
"sourcedock.dev/petrbalvin/tensor/internal/core"
|
||||
)
|
||||
|
||||
// The boolean fixture bool_enum.h5 under testdata/h5 was written by
|
||||
// hand against the HDF5 1.8 specification, in the shape the reference
|
||||
// library writes for an old-style group with a symbol table, and it
|
||||
// validates against the reference library's own tools: h5dump reads
|
||||
// the dataset back as an H5T_ENUM over H5T_STD_U8LE with the members
|
||||
// FALSE = 0 and TRUE = 1 and the values TRUE, FALSE, TRUE, TRUE,
|
||||
// FALSE. Its layout deliberately differs from what this package's
|
||||
// writer produces: the messages of the dataset's object header sit in
|
||||
// the order datatype, fill value, dataspace, layout, the dataspace
|
||||
// carries no maximum dimensions and the fill value message declares a
|
||||
// defined zero fill, so the pins below prove the reader's tolerance of
|
||||
// a foreign layout rather than a round trip of its own bytes.
|
||||
|
||||
// TestLoadHDF5ForeignBoolFixture pins the read side: the hand-crafted
|
||||
// boolean enumeration of a foreign layout lands the dataset "flags"
|
||||
// with the core Bool dtype and the values the reference library reads
|
||||
// from the same bytes.
|
||||
func TestLoadHDF5ForeignBoolFixture(t *testing.T) {
|
||||
sets, err := LoadHDF5(h5Fixture(t, "bool_enum.h5"))
|
||||
if err != nil {
|
||||
t.Fatalf("LoadHDF5: %v", err)
|
||||
}
|
||||
if len(sets) != 1 {
|
||||
t.Fatalf("datasets = %d, want 1", len(sets))
|
||||
}
|
||||
d := sets[0]
|
||||
if d.Path != "/flags" {
|
||||
t.Fatalf("path = %q, want /flags", d.Path)
|
||||
}
|
||||
if s := d.Shape; len(s) != 1 || s[0] != 5 {
|
||||
t.Fatalf("shape = %v, want [5]", s)
|
||||
}
|
||||
if dt := d.Values.Dtype(); dt != core.Bool {
|
||||
t.Fatalf("dtype = %s, want bool", dt)
|
||||
}
|
||||
want := []bool{true, false, true, true, false}
|
||||
if got := d.Values.RawBools()[:5]; !slices.Equal(got, want) {
|
||||
t.Fatalf("values = %v, want %v", got, want)
|
||||
}
|
||||
}
|
||||
|
||||
// TestSaveHDF5BoolConvention pins the write side against the same
|
||||
// logical values: SaveHDF5 stores them through the boolean enumeration
|
||||
// convention and LoadHDF5 reads them back identically. The pin holds
|
||||
// the two-sided agreement on the convention, not a byte match with the
|
||||
// foreign fixture, whose layout the writer need not reproduce.
|
||||
func TestSaveHDF5BoolConvention(t *testing.T) {
|
||||
want := []bool{true, false, true, true, false}
|
||||
values, err := core.FromBools(want, 5)
|
||||
if err != nil {
|
||||
t.Fatalf("FromBools: %v", err)
|
||||
}
|
||||
path := filepath.Join(t.TempDir(), "bool_convention.h5")
|
||||
if err := SaveHDF5(path, []HDF5Dataset{{Path: "/flags", Shape: []int{5}, Values: values}}, nil); err != nil {
|
||||
t.Fatalf("SaveHDF5: %v", err)
|
||||
}
|
||||
sets, err := LoadHDF5(path)
|
||||
if err != nil {
|
||||
t.Fatalf("LoadHDF5: %v", err)
|
||||
}
|
||||
if len(sets) != 1 {
|
||||
t.Fatalf("datasets = %d, want 1", len(sets))
|
||||
}
|
||||
d := sets[0]
|
||||
if d.Path != "/flags" {
|
||||
t.Fatalf("path = %q, want /flags", d.Path)
|
||||
}
|
||||
if s := d.Shape; len(s) != 1 || s[0] != 5 {
|
||||
t.Fatalf("shape = %v, want [5]", s)
|
||||
}
|
||||
if dt := d.Values.Dtype(); dt != core.Bool {
|
||||
t.Fatalf("dtype = %s, want bool", dt)
|
||||
}
|
||||
if got := d.Values.RawBools()[:5]; !slices.Equal(got, want) {
|
||||
t.Fatalf("values = %v, want %v", got, want)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,161 @@
|
||||
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
|
||||
// SPDX-License-Identifier: MIT
|
||||
|
||||
package io
|
||||
|
||||
import (
|
||||
"encoding/binary"
|
||||
"math"
|
||||
"strings"
|
||||
"testing"
|
||||
)
|
||||
|
||||
// Regression pins: HDF5 headers the reference library writes
|
||||
// once they outgrow their first block, byte orders and unallocated
|
||||
// storage the reader used to answer silently, and the CSV byte order
|
||||
// mark every spreadsheet writes.
|
||||
|
||||
// h5Link renders a continuation link message body: the block address
|
||||
// and its length, eight bytes each in the hostile layout.
|
||||
func h5Link(addr, length uint64) []byte {
|
||||
b := make([]byte, 16)
|
||||
binary.LittleEndian.PutUint64(b, addr)
|
||||
binary.LittleEndian.PutUint64(b[8:], length)
|
||||
return b
|
||||
}
|
||||
|
||||
// h5ContinuationBlock writes a flat message-list block at off and
|
||||
// returns the end offset, everything eight-aligned.
|
||||
func h5ContinuationBlock(f []byte, off int, msgs ...h5Msg) int {
|
||||
for _, m := range msgs {
|
||||
binary.LittleEndian.PutUint16(f[off:], m.typ)
|
||||
binary.LittleEndian.PutUint16(f[off+2:], uint16(len(m.body)))
|
||||
copy(f[off+8:], m.body)
|
||||
off += alignUp(8+len(m.body), 8)
|
||||
}
|
||||
return off
|
||||
}
|
||||
|
||||
// TestLoadHDF5ChainedContinuation: the header walk followed exactly
|
||||
// one continuation block and dropped every message of the second and
|
||||
// later ones, so a legal attribute-rich file lost its datasets. Here
|
||||
// the dataspace sits in the header, the datatype in the first
|
||||
// continuation block and the layout in the second: only a walk that
|
||||
// follows the whole chain can assemble the dataset.
|
||||
func TestLoadHDF5ChainedContinuation(t *testing.T) {
|
||||
const blockA, blockB, dataAt, n = 448, 640, 800, 1024
|
||||
f := h5HostileFile(n)
|
||||
end := h5ObjectHeader(f, 96,
|
||||
h5Msg{hdf5MsgDataspace, h5Dataspace(2)},
|
||||
h5Msg{hdf5MsgContinuation, h5Link(blockA, 112)},
|
||||
)
|
||||
if end > blockA {
|
||||
t.Fatalf("the header runs to %d, past the first block at %d", end, blockA)
|
||||
}
|
||||
endA := h5ContinuationBlock(f, blockA,
|
||||
h5Msg{hdf5MsgDatatype, h5FloatType(8)},
|
||||
h5Msg{hdf5MsgContinuation, h5Link(blockB, 88)},
|
||||
)
|
||||
if endA > blockB {
|
||||
t.Fatalf("block A runs to %d, past block B at %d", endA, blockB)
|
||||
}
|
||||
h5ContinuationBlock(f, blockB,
|
||||
h5Msg{hdf5MsgDataLayout, h5ContiguousLayout(dataAt, 16)},
|
||||
)
|
||||
binary.LittleEndian.PutUint64(f[dataAt:], math.Float64bits(1.5))
|
||||
binary.LittleEndian.PutUint64(f[dataAt+8:], math.Float64bits(-2.5))
|
||||
sets, err := LoadHDF5(writeHostile(t, "chained.h5", f))
|
||||
if err != nil {
|
||||
t.Fatalf("LoadHDF5: %v", err)
|
||||
}
|
||||
if len(sets) != 1 {
|
||||
t.Fatalf("datasets = %d, want 1", len(sets))
|
||||
}
|
||||
v0 := sets[0].Values.FloatAt(0)
|
||||
if v0 != 1.5 || sets[0].Values.FloatAt(1) != -2.5 {
|
||||
t.Fatalf("values = %v, %v, want 1.5 and -2.5", v0, sets[0].Values.FloatAt(1))
|
||||
}
|
||||
}
|
||||
|
||||
// TestLoadHDF5ContinuationCycle: a continuation block listing itself
|
||||
// must be an error, not an infinite walk.
|
||||
func TestLoadHDF5ContinuationCycle(t *testing.T) {
|
||||
const blockA, n = 448, 640
|
||||
f := h5HostileFile(n)
|
||||
h5ObjectHeader(f, 96,
|
||||
h5Msg{hdf5MsgDataspace, h5Dataspace(2)},
|
||||
h5Msg{hdf5MsgContinuation, h5Link(blockA, 40)},
|
||||
)
|
||||
// The block's only entry is a link back to itself.
|
||||
h5ContinuationBlock(f, blockA,
|
||||
h5Msg{hdf5MsgContinuation, h5Link(blockA, 40)},
|
||||
)
|
||||
if _, err := LoadHDF5(writeHostile(t, "cycle.h5", f)); err == nil || !strings.Contains(err.Error(), "twice") {
|
||||
t.Fatalf("a self-referencing continuation block: err = %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
// TestLoadHDF5BigEndianRefusal: the byte-order bit of the datatype was
|
||||
// parsed and never checked, so a big-endian dataset decoded as
|
||||
// byte-swapped noise with no error.
|
||||
func TestLoadHDF5BigEndianRefusal(t *testing.T) {
|
||||
const dataAt, n = 448, 512
|
||||
f := h5HostileFile(n)
|
||||
beType := h5FloatType(8)
|
||||
beType[1] = 0x01 // class bit field: bit 0 set means big-endian
|
||||
h5ObjectHeader(f, 96,
|
||||
h5Msg{hdf5MsgDataspace, h5Dataspace(2)},
|
||||
h5Msg{hdf5MsgDatatype, beType},
|
||||
h5Msg{hdf5MsgDataLayout, h5ContiguousLayout(dataAt, 16)},
|
||||
)
|
||||
if _, err := LoadHDF5(writeHostile(t, "bigendian.h5", f)); err == nil || !strings.Contains(err.Error(), "big-endian") {
|
||||
t.Fatalf("a big-endian dataset: err = %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
// TestLoadHDF5UnallocatedContiguous: the undefined storage address on
|
||||
// a contiguous dataset refused the whole file; an empty dataset there
|
||||
// is legal (nothing was ever allocated) and must load as empty, while
|
||||
// a non-empty one names what is missing.
|
||||
func TestLoadHDF5UnallocatedContiguous(t *testing.T) {
|
||||
t.Run("empty", func(t *testing.T) {
|
||||
f := h5HostileFile(512)
|
||||
h5ObjectHeader(f, 96,
|
||||
h5Msg{hdf5MsgDataspace, h5Dataspace(0)},
|
||||
h5Msg{hdf5MsgDatatype, h5FloatType(8)},
|
||||
h5Msg{hdf5MsgDataLayout, h5ContiguousLayout(math.MaxUint64, 0)},
|
||||
)
|
||||
sets, err := LoadHDF5(writeHostile(t, "empty.h5", f))
|
||||
if err != nil {
|
||||
t.Fatalf("an empty unallocated dataset: %v", err)
|
||||
}
|
||||
if len(sets) != 1 || sets[0].Values.Len() != 0 {
|
||||
t.Fatalf("datasets = %d, want one empty dataset", len(sets))
|
||||
}
|
||||
})
|
||||
t.Run("non-empty", func(t *testing.T) {
|
||||
f := h5HostileFile(512)
|
||||
h5ObjectHeader(f, 96,
|
||||
h5Msg{hdf5MsgDataspace, h5Dataspace(2)},
|
||||
h5Msg{hdf5MsgDatatype, h5FloatType(8)},
|
||||
h5Msg{hdf5MsgDataLayout, h5ContiguousLayout(math.MaxUint64, 16)},
|
||||
)
|
||||
if _, err := LoadHDF5(writeHostile(t, "noalloc.h5", f)); err == nil || !strings.Contains(err.Error(), "never allocated") {
|
||||
t.Fatalf("a non-empty unallocated dataset: err = %v", err)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
// TestLoadCSVSkipBOM: a leading UTF-8 byte order mark used to glue
|
||||
// itself onto the first field and fail the whole load with a strconv
|
||||
// error.
|
||||
func TestLoadCSVSkipBOM(t *testing.T) {
|
||||
const in = "\xEF\xBB\xBF1.5,2.5\n3,4\n"
|
||||
a, err := LoadCSVReader(strings.NewReader(in), false)
|
||||
if err != nil {
|
||||
t.Fatalf("LoadCSVReader with a BOM: %v", err)
|
||||
}
|
||||
if a.FloatAt(0) != 1.5 || a.FloatAt(1) != 2.5 || a.FloatAt(2) != 3 || a.FloatAt(3) != 4 {
|
||||
t.Fatalf("values = %v", a)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,527 @@
|
||||
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
|
||||
// SPDX-License-Identifier: MIT
|
||||
|
||||
package io
|
||||
|
||||
import (
|
||||
"encoding/binary"
|
||||
"math"
|
||||
"slices"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
)
|
||||
|
||||
// Regression pins: continuation links the reader sliced past the
|
||||
// end of, a local heap header sized by the wrong field, the diamond of
|
||||
// hard links that multiplied the walk, the read budget that only
|
||||
// counted value bytes, the B-tree sizes pinned to eight-byte addresses,
|
||||
// negative FITS column counts, the NAXIS prefix that swallowed user
|
||||
// keywords, and the version 2 header features no test exercised.
|
||||
|
||||
// h5SizedHostileFile returns an n-byte HDF5 file with the signature and
|
||||
// a version 0 superblock of the given address and length sizes whose
|
||||
// root object header sits at rootAt. h5HostileFile is the 8/8 case;
|
||||
// the hostile files here need the others.
|
||||
func h5SizedHostileFile(n, offSize, lenSize int, rootAt uint64) []byte {
|
||||
f := make([]byte, n)
|
||||
copy(f, hdf5Magic)
|
||||
f[8] = 0 // superblock version 0
|
||||
f[13] = byte(offSize)
|
||||
f[14] = byte(lenSize)
|
||||
// Four addresses of offSize bytes from 24: base, free space, end of
|
||||
// file, driver information.
|
||||
putAddr := func(at int, v uint64) {
|
||||
switch offSize {
|
||||
case 4:
|
||||
binary.LittleEndian.PutUint32(f[at:], uint32(v))
|
||||
case 8:
|
||||
binary.LittleEndian.PutUint64(f[at:], v)
|
||||
}
|
||||
}
|
||||
putAddr(24, 0)
|
||||
putAddr(28, math.MaxUint64)
|
||||
putAddr(32, uint64(n))
|
||||
putAddr(36, math.MaxUint64)
|
||||
// The root symbol table entry: a link name offset of the length
|
||||
// size, then the object header address of the address size.
|
||||
entry := 24 + 4*offSize
|
||||
putAddr(entry+lenSize, rootAt)
|
||||
return f
|
||||
}
|
||||
|
||||
// h5HardLink renders a version 1 hard-link message body: a one-byte
|
||||
// name length, the name and the object header address.
|
||||
func h5HardLink(name string, addr uint64) []byte {
|
||||
b := append([]byte{1, 0, byte(len(name))}, name...)
|
||||
return binary.LittleEndian.AppendUint64(b, addr)
|
||||
}
|
||||
|
||||
// h5V3Superblock returns an n-byte file with a version 3 superblock
|
||||
// (eight-byte addresses and lengths, the whole block under a lookup3
|
||||
// checksum) whose root object header sits at rootAt.
|
||||
func h5V3Superblock(n int, rootAt uint64) []byte {
|
||||
f := make([]byte, n)
|
||||
copy(f, hdf5Magic)
|
||||
f[8] = 3
|
||||
f[9] = 8 // address size
|
||||
f[10] = 8 // length size
|
||||
f[11] = 0 // consistency flags
|
||||
binary.LittleEndian.PutUint64(f[12:], 0) // base address
|
||||
binary.LittleEndian.PutUint64(f[20:], math.MaxUint64) // no extension
|
||||
binary.LittleEndian.PutUint64(f[28:], uint64(n)) // end of file
|
||||
binary.LittleEndian.PutUint64(f[36:], rootAt)
|
||||
binary.LittleEndian.PutUint32(f[44:], hdf5Lookup3(f[:44]))
|
||||
return f
|
||||
}
|
||||
|
||||
// h5WriteOHDRv2 writes a version 2 object header at at whose first
|
||||
// message region is chunk, under the given flag byte, computing and
|
||||
// storing the lookup3 checksum, and returns the offset just past the
|
||||
// header. Flag-selected prefix fields are written as zeros.
|
||||
func h5WriteOHDRv2(f []byte, at int, flags byte, chunk []byte) int {
|
||||
copy(f[at:], hdf5ObjHdr2)
|
||||
f[at+4] = 2
|
||||
f[at+5] = flags
|
||||
p := at + 6
|
||||
if flags&0x20 != 0 {
|
||||
p += 16 // access, modification, change and birth times
|
||||
}
|
||||
if flags&0x10 != 0 {
|
||||
p += 4 // max compact and min dense attribute counts
|
||||
}
|
||||
width := 1 << (flags & 0x03)
|
||||
switch width {
|
||||
case 1:
|
||||
f[p] = byte(len(chunk))
|
||||
case 2:
|
||||
binary.LittleEndian.PutUint16(f[p:], uint16(len(chunk)))
|
||||
case 4:
|
||||
binary.LittleEndian.PutUint32(f[p:], uint32(len(chunk)))
|
||||
case 8:
|
||||
binary.LittleEndian.PutUint64(f[p:], uint64(len(chunk)))
|
||||
}
|
||||
p += width
|
||||
copy(f[p:], chunk)
|
||||
p += len(chunk)
|
||||
binary.LittleEndian.PutUint32(f[p:], hdf5Lookup3(f[at:p]))
|
||||
return p + 4
|
||||
}
|
||||
|
||||
// h5V2Msg renders one version 2 object header message: a type byte, a
|
||||
// two-byte size and a flag byte, widened by a two-byte creation order
|
||||
// when the header tracks it, then the body.
|
||||
func h5V2Msg(typ byte, order uint16, body []byte, ordered bool) []byte {
|
||||
head := 4
|
||||
if ordered {
|
||||
head = 6
|
||||
}
|
||||
m := make([]byte, head+len(body))
|
||||
m[0] = typ
|
||||
binary.LittleEndian.PutUint16(m[1:], uint16(len(body)))
|
||||
if ordered {
|
||||
binary.LittleEndian.PutUint16(m[4:], order)
|
||||
}
|
||||
copy(m[head:], body)
|
||||
return m
|
||||
}
|
||||
|
||||
// h5V2Dataspace renders a version 2 dataspace message body: rank one,
|
||||
// the given extent in eight bytes.
|
||||
func h5V2Dataspace(dim uint64) []byte {
|
||||
b := make([]byte, 12)
|
||||
b[0] = 2 // version
|
||||
b[1] = 1 // rank
|
||||
binary.LittleEndian.PutUint64(b[4:], dim)
|
||||
return b
|
||||
}
|
||||
|
||||
// TestLoadHDF5ShortV1Continuation pins a minimal hostile
|
||||
// file: a continuation block that ends right after its eight-byte
|
||||
// message header carries a zero-size continuation message, and the
|
||||
// reader used to slice the link's offset and length out of bytes past
|
||||
// the block (a [16:8] panic out of LoadHDF5). It must be a named
|
||||
// error.
|
||||
func TestLoadHDF5ShortV1Continuation(t *testing.T) {
|
||||
f := h5HostileFile(512)
|
||||
h5ObjectHeader(f, 96,
|
||||
h5Msg{hdf5MsgContinuation, h5Link(136, 8)},
|
||||
)
|
||||
// The block: exactly eight bytes, one continuation header, no body.
|
||||
binary.LittleEndian.PutUint16(f[136:], hdf5MsgContinuation)
|
||||
binary.LittleEndian.PutUint16(f[138:], 0)
|
||||
_, err := LoadHDF5(writeHostile(t, "v1short.h5", f))
|
||||
if err == nil {
|
||||
t.Fatal("LoadHDF5 accepted a continuation message with no link body")
|
||||
}
|
||||
if !strings.Contains(err.Error(), "shorter than an offset and a length") {
|
||||
t.Fatalf("error = %v, want the short-link refusal", err)
|
||||
}
|
||||
}
|
||||
|
||||
// TestLoadHDF5ShortV2Continuation pins the version 2 twin: a header
|
||||
// whose whole message region is one continuation message of size zero
|
||||
// panicked the same way ([12:4]), because the stream read the link's
|
||||
// offset and length past the region.
|
||||
func TestLoadHDF5ShortV2Continuation(t *testing.T) {
|
||||
f := h5V3Superblock(512, 96)
|
||||
chunk := []byte{hdf5MsgContinuation, 0, 0, 0} // type 16, size 0, flags 0
|
||||
h5WriteOHDRv2(f, 96, 0, chunk)
|
||||
_, err := LoadHDF5(writeHostile(t, "v2short.h5", f))
|
||||
if err == nil {
|
||||
t.Fatal("LoadHDF5 accepted a version 2 continuation message with no link body")
|
||||
}
|
||||
if !strings.Contains(err.Error(), "shorter than an offset and a length") {
|
||||
t.Fatalf("error = %v, want the short-link refusal", err)
|
||||
}
|
||||
}
|
||||
|
||||
// TestLoadHDF5LocalHeapMixedSizes pins the local heap header size: the
|
||||
// header holds two length-size fields and then one address-size field,
|
||||
// but the reader sized it with three address-size fields, so a 4/8 file
|
||||
// (twenty-byte header) was sliced at [24:] and panicked. A heap whose
|
||||
// header lies about nothing must still be refused cleanly when what it
|
||||
// points at is absent.
|
||||
func TestLoadHDF5LocalHeapMixedSizes(t *testing.T) {
|
||||
f := h5SizedHostileFile(512, 4, 8, 96)
|
||||
body := make([]byte, 2*4)
|
||||
binary.LittleEndian.PutUint32(body, math.MaxUint32) // B-tree: undefined
|
||||
binary.LittleEndian.PutUint32(body[4:], 256) // local heap at 256
|
||||
h5ObjectHeader(f, 96, h5Msg{hdf5MsgSymbolTable, body})
|
||||
copy(f[256:], hdf5LocalHeap)
|
||||
f[260] = 1 // version
|
||||
_, err := LoadHDF5(writeHostile(t, "heap48.h5", f))
|
||||
if err != nil && strings.Contains(err.Error(), "runtime error") {
|
||||
t.Fatalf("the 4/8 local heap panicked: %v", err)
|
||||
}
|
||||
_ = err // any named refusal is fine; the panic is the defect
|
||||
}
|
||||
|
||||
// TestWalkDiamondReadsOnce pins the diamond: two names on one group
|
||||
// used to walk the group twice, and a diamond of depth d walked it 2^d
|
||||
// times, which stopped the reader for hours on a file of a kilobyte.
|
||||
// The object is read once, under the first path the traversal reaches,
|
||||
// and a link that closes a cycle along the current path is still an
|
||||
// error.
|
||||
func TestWalkDiamondReadsOnce(t *testing.T) {
|
||||
guard := time.AfterFunc(20*time.Second, func() { panic("diamond walk did not return") })
|
||||
defer guard.Stop()
|
||||
|
||||
const childAt, datasetAt, dataAt = 256, 384, 512
|
||||
f := h5HostileFile(576)
|
||||
h5ObjectHeader(f, 96,
|
||||
h5Msg{hdf5MsgLink, h5HardLink("a", childAt)},
|
||||
h5Msg{hdf5MsgLink, h5HardLink("b", childAt)},
|
||||
)
|
||||
h5ObjectHeader(f, childAt,
|
||||
h5Msg{hdf5MsgLink, h5HardLink("d", datasetAt)},
|
||||
)
|
||||
h5ObjectHeader(f, datasetAt,
|
||||
h5Msg{hdf5MsgDataspace, h5Dataspace(1)},
|
||||
h5Msg{hdf5MsgDatatype, h5FloatType(8)},
|
||||
h5Msg{hdf5MsgDataLayout, h5ContiguousLayout(dataAt, 8)},
|
||||
)
|
||||
binary.LittleEndian.PutUint64(f[dataAt:], math.Float64bits(2.5))
|
||||
sets, err := LoadHDF5(writeHostile(t, "diamond.h5", f))
|
||||
if err != nil {
|
||||
t.Fatalf("LoadHDF5 on a diamond of hard links: %v", err)
|
||||
}
|
||||
if len(sets) != 1 {
|
||||
t.Fatalf("datasets = %d, want the single dataset once", len(sets))
|
||||
}
|
||||
// Deterministic: the links of a group are visited in file order, so
|
||||
// the first path wins and the listing is stable.
|
||||
if sets[0].Path != "/a/d" {
|
||||
t.Fatalf("path = %q, want /a/d, the first path to the object", sets[0].Path)
|
||||
}
|
||||
if got := sets[0].Values.FloatAt(0); got != 2.5 {
|
||||
t.Fatalf("value = %v, want 2.5", got)
|
||||
}
|
||||
t.Run("cycle", func(t *testing.T) {
|
||||
// The skip must never swallow a cycle: an object that closes a
|
||||
// loop along the current path is refused, not skipped.
|
||||
if _, err := LoadHDF5(writeHostile(t, "diamond-cycle.h5", selfLink())); err == nil || !strings.Contains(err.Error(), "hard-link cycle") {
|
||||
t.Fatalf("a hard-link cycle: err = %v, want the cycle refusal", err)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
// TestHDF5BudgetChargesDatasetEnvelope pins the aggregate budget: it
|
||||
// used to deduct only the value bytes, so millions of empty datasets
|
||||
// would allocate their result structures, paths and attribute maps
|
||||
// without ever touching the budget. The walk charges a fixed envelope
|
||||
// plus the dataset's path beside the values, so a budget below one
|
||||
// envelope refuses even a file of empty datasets.
|
||||
func TestHDF5BudgetChargesDatasetEnvelope(t *testing.T) {
|
||||
const dataAt, n = 448, 512
|
||||
raw := h5HostileFile(n)
|
||||
h5ObjectHeader(raw, 96,
|
||||
h5Msg{hdf5MsgDataspace, h5Dataspace(1)},
|
||||
h5Msg{hdf5MsgDatatype, h5FloatType(8)},
|
||||
h5Msg{hdf5MsgDataLayout, h5ContiguousLayout(dataAt, 8)},
|
||||
)
|
||||
binary.LittleEndian.PutUint64(raw[dataAt:], math.Float64bits(2.5))
|
||||
// The file itself is legal: at the default budget it loads.
|
||||
sets, err := LoadHDF5(writeHostile(t, "envelope.h5", raw))
|
||||
if err != nil || len(sets) != 1 {
|
||||
t.Fatalf("LoadHDF5 at the default budget: %d datasets, err = %v", len(sets), err)
|
||||
}
|
||||
// Below one envelope the same file refuses: the envelope plus the
|
||||
// path is charged before any allocation happens.
|
||||
f, err := newHDF5File(raw)
|
||||
if err != nil {
|
||||
t.Fatalf("newHDF5File: %v", err)
|
||||
}
|
||||
st := newHDF5WalkState()
|
||||
st.budget = hdf5DatasetEnvelope / 2
|
||||
var out []HDF5Dataset
|
||||
if err := f.walk(f.rootAddress, st, &out); err == nil || !strings.Contains(err.Error(), "budget") {
|
||||
t.Fatalf("walk with a budget below one envelope: err = %v, want the budget refusal", err)
|
||||
}
|
||||
}
|
||||
|
||||
// TestLoadHDF5SymbolTable4of4 pins the four-byte-address layout of the
|
||||
// group structures: the version 1 B-tree header is 8+2*offSize wide
|
||||
// (not 24), a symbol node entry is lenSize+offSize+24 (not 40) and the
|
||||
// local heap header is 8+2*lenSize+offSize. A legal 4/4 file with a
|
||||
// symbol-table group of two links used to die with "no symbol table
|
||||
// node at 0"; the 8/8 fixtures are untouched by the arithmetic.
|
||||
func TestLoadHDF5SymbolTable4of4(t *testing.T) {
|
||||
const (
|
||||
treeAt = 200
|
||||
snodAt = 240
|
||||
heapAt = 320
|
||||
segAt = 352
|
||||
objA = 384
|
||||
objB = 512
|
||||
dataA = 640
|
||||
dataB = 648
|
||||
)
|
||||
f := h5SizedHostileFile(1024, 4, 4, 96)
|
||||
body := make([]byte, 8)
|
||||
binary.LittleEndian.PutUint32(body, treeAt)
|
||||
binary.LittleEndian.PutUint32(body[4:], heapAt)
|
||||
h5ObjectHeader(f, 96, h5Msg{hdf5MsgSymbolTable, body})
|
||||
|
||||
// The group B-tree: a sixteen-byte header (signature, type, level,
|
||||
// one entry, two sibling addresses), key0 of four bytes, the child
|
||||
// address, one trailing key.
|
||||
copy(f[treeAt:], hdf5Tree)
|
||||
f[treeAt+4] = 0 // group
|
||||
f[treeAt+5] = 0 // leaf
|
||||
binary.LittleEndian.PutUint16(f[treeAt+6:], 1)
|
||||
binary.LittleEndian.PutUint32(f[treeAt+8:], math.MaxUint32)
|
||||
binary.LittleEndian.PutUint32(f[treeAt+12:], math.MaxUint32)
|
||||
binary.LittleEndian.PutUint32(f[treeAt+16:], 0) // key0
|
||||
binary.LittleEndian.PutUint32(f[treeAt+20:], snodAt) // child
|
||||
binary.LittleEndian.PutUint32(f[treeAt+24:], 4) // trailing key
|
||||
|
||||
// The symbol node: two entries of 32 bytes (heap offset, object
|
||||
// address, cache type, reserved, scratch pad).
|
||||
copy(f[snodAt:], hdf5SymbolNode)
|
||||
f[snodAt+4] = 1
|
||||
binary.LittleEndian.PutUint16(f[snodAt+6:], 2)
|
||||
binary.LittleEndian.PutUint32(f[snodAt+8:], 0) // "a" at heap offset 0
|
||||
binary.LittleEndian.PutUint32(f[snodAt+12:], objA)
|
||||
binary.LittleEndian.PutUint32(f[snodAt+40:], 2) // "b" at heap offset 2
|
||||
binary.LittleEndian.PutUint32(f[snodAt+44:], objB)
|
||||
|
||||
// The local heap: a twenty-byte header (signature, version, data
|
||||
// segment size, free-list head, data segment address).
|
||||
copy(f[heapAt:], hdf5LocalHeap)
|
||||
f[heapAt+4] = 1
|
||||
binary.LittleEndian.PutUint32(f[heapAt+8:], 8)
|
||||
binary.LittleEndian.PutUint32(f[heapAt+12:], math.MaxUint32)
|
||||
binary.LittleEndian.PutUint32(f[heapAt+16:], segAt)
|
||||
copy(f[segAt:], "a\x00b\x00\x00\x00\x00\x00")
|
||||
|
||||
h5ObjectHeader(f, objA,
|
||||
h5Msg{hdf5MsgDataspace, h5Dataspace(1)},
|
||||
h5Msg{hdf5MsgDatatype, h5FloatType(8)},
|
||||
h5Msg{hdf5MsgDataLayout, h5ContiguousLayout(dataA, 8)},
|
||||
)
|
||||
h5ObjectHeader(f, objB,
|
||||
h5Msg{hdf5MsgDataspace, h5Dataspace(1)},
|
||||
h5Msg{hdf5MsgDatatype, h5FloatType(8)},
|
||||
h5Msg{hdf5MsgDataLayout, h5ContiguousLayout(dataB, 8)},
|
||||
)
|
||||
binary.LittleEndian.PutUint64(f[dataA:], math.Float64bits(4.5))
|
||||
binary.LittleEndian.PutUint64(f[dataB:], math.Float64bits(-2.5))
|
||||
|
||||
sets, err := LoadHDF5(writeHostile(t, "snod44.h5", f))
|
||||
if err != nil {
|
||||
t.Fatalf("LoadHDF5 refused a legal 4/4 symbol-table file: %v", err)
|
||||
}
|
||||
if len(sets) != 2 {
|
||||
t.Fatalf("datasets = %d, want 2", len(sets))
|
||||
}
|
||||
if sets[0].Path != "/a" || sets[1].Path != "/b" {
|
||||
t.Fatalf("paths = %q, %q, want /a and /b", sets[0].Path, sets[1].Path)
|
||||
}
|
||||
if got := sets[0].Values.FloatAt(0); got != 4.5 {
|
||||
t.Fatalf("/a = %v, want 4.5", got)
|
||||
}
|
||||
if got := sets[1].Values.FloatAt(0); got != -2.5 {
|
||||
t.Fatalf("/b = %v, want -2.5", got)
|
||||
}
|
||||
}
|
||||
|
||||
// TestLoadHDF5ContinuationChainDepth pins the version 1 depth guard: a
|
||||
// chain of continuation blocks one longer than the cap must be refused
|
||||
// like the version 2 walk refuses its own, not walked to the end.
|
||||
func TestLoadHDF5ContinuationChainDepth(t *testing.T) {
|
||||
const first = 224
|
||||
blocks := hdf5MaxHeaderBlocks + 2
|
||||
f := h5HostileFile(first + blocks*24)
|
||||
h5ObjectHeader(f, 96,
|
||||
h5Msg{hdf5MsgContinuation, h5Link(first, 24)},
|
||||
)
|
||||
for k := range blocks {
|
||||
at := first + k*24
|
||||
binary.LittleEndian.PutUint16(f[at:], hdf5MsgContinuation)
|
||||
if k == blocks-1 {
|
||||
// The tail block carries nothing: the walk must be refused
|
||||
// on reaching it, not on reading it.
|
||||
binary.LittleEndian.PutUint16(f[at+2:], 0)
|
||||
continue
|
||||
}
|
||||
binary.LittleEndian.PutUint16(f[at+2:], 16)
|
||||
binary.LittleEndian.PutUint64(f[at+8:], uint64(at+24))
|
||||
binary.LittleEndian.PutUint64(f[at+16:], 24)
|
||||
}
|
||||
_, err := LoadHDF5(writeHostile(t, "chain.h5", f))
|
||||
if err == nil {
|
||||
t.Fatalf("LoadHDF5 walked a chain of %d continuation blocks", blocks)
|
||||
}
|
||||
if !strings.Contains(err.Error(), "continuation blocks") {
|
||||
t.Fatalf("error = %v, want the chain-depth refusal", err)
|
||||
}
|
||||
}
|
||||
|
||||
// TestLoadHDF5V2HeaderFlags covers the version 2 header features the
|
||||
// reference fixture does not carry: the four times (0x20), the attribute
|
||||
// counts (0x10) and creation order tracking (0x04, which widens every
|
||||
// message by its order field). A header with all three set must read
|
||||
// like any other.
|
||||
func TestLoadHDF5V2HeaderFlags(t *testing.T) {
|
||||
const dataAt = 320
|
||||
f := h5V3Superblock(512, 96)
|
||||
msgs := slices.Concat(
|
||||
h5V2Msg(hdf5MsgDataspace, 1, h5V2Dataspace(2), true),
|
||||
h5V2Msg(hdf5MsgDatatype, 2, h5FloatType(8), true),
|
||||
h5V2Msg(hdf5MsgDataLayout, 3, h5ContiguousLayout(dataAt, 16), true),
|
||||
)
|
||||
h5WriteOHDRv2(f, 96, 0x34, msgs)
|
||||
binary.LittleEndian.PutUint64(f[dataAt:], math.Float64bits(1.5))
|
||||
binary.LittleEndian.PutUint64(f[dataAt+8:], math.Float64bits(-2.5))
|
||||
|
||||
sets, err := LoadHDF5(writeHostile(t, "v2flags.h5", f))
|
||||
if err != nil {
|
||||
t.Fatalf("LoadHDF5 refused a flagged version 2 header: %v", err)
|
||||
}
|
||||
if len(sets) != 1 || sets[0].Path != "/" {
|
||||
t.Fatalf("datasets = %v, want one dataset at /", sets)
|
||||
}
|
||||
if got := sets[0].Values.FloatAt(1); got != -2.5 {
|
||||
t.Fatalf("value = %v, want -2.5", got)
|
||||
}
|
||||
}
|
||||
|
||||
// TestLoadHDF5V2ContinuationChecksum covers the happy path of a version
|
||||
// 2 continuation: the dataset's messages live in an OCHK block whose
|
||||
// lookup3 checksum covers the signature and the messages alike, and a
|
||||
// correct checksum must be accepted (the hostile files above only ever
|
||||
// see a broken one).
|
||||
func TestLoadHDF5V2ContinuationChecksum(t *testing.T) {
|
||||
const blockAt, dataAt = 224, 320
|
||||
f := h5V3Superblock(512, 96)
|
||||
msgs := slices.Concat(
|
||||
h5V2Msg(hdf5MsgDataspace, 0, h5V2Dataspace(2), false),
|
||||
h5V2Msg(hdf5MsgDatatype, 0, h5FloatType(8), false),
|
||||
h5V2Msg(hdf5MsgDataLayout, 0, h5ContiguousLayout(dataAt, 16), false),
|
||||
)
|
||||
block := append([]byte{}, hdf5Chunk2...)
|
||||
block = append(block, msgs...)
|
||||
block = binary.LittleEndian.AppendUint32(block, hdf5Lookup3(block))
|
||||
copy(f[blockAt:], block)
|
||||
|
||||
link := binary.LittleEndian.AppendUint64(
|
||||
binary.LittleEndian.AppendUint64([]byte{}, blockAt), uint64(len(block)))
|
||||
chunk := make([]byte, 4+len(link))
|
||||
chunk[0] = hdf5MsgContinuation
|
||||
binary.LittleEndian.PutUint16(chunk[1:], uint16(len(link)))
|
||||
copy(chunk[4:], link)
|
||||
h5WriteOHDRv2(f, 96, 0, chunk)
|
||||
|
||||
binary.LittleEndian.PutUint64(f[dataAt:], math.Float64bits(1.5))
|
||||
binary.LittleEndian.PutUint64(f[dataAt+8:], math.Float64bits(-2.5))
|
||||
|
||||
sets, err := LoadHDF5(writeHostile(t, "v2cont-ok.h5", f))
|
||||
if err != nil {
|
||||
t.Fatalf("LoadHDF5 refused a checksummed version 2 continuation: %v", err)
|
||||
}
|
||||
if len(sets) != 1 || sets[0].Path != "/" {
|
||||
t.Fatalf("datasets = %v, want one dataset at /", sets)
|
||||
}
|
||||
if got := sets[0].Values.FloatAt(0); got != 1.5 {
|
||||
t.Fatalf("value = %v, want 1.5", got)
|
||||
}
|
||||
}
|
||||
|
||||
// TestLoadFITSTableNegativeTFIELDS pins the column count: a negative
|
||||
// TFIELDS used to size the per-column slices with a negative length and
|
||||
// panicked, in the binary and the ASCII branch alike. It is refused by
|
||||
// name in both.
|
||||
func TestLoadFITSTableNegativeTFIELDS(t *testing.T) {
|
||||
for _, kind := range []string{"BINTABLE", "TABLE"} {
|
||||
t.Run(kind, func(t *testing.T) {
|
||||
hdr := cardBlock(
|
||||
card("XTENSION= '"+kind+"'"),
|
||||
card("BITPIX = 8"),
|
||||
card("NAXIS = 2"),
|
||||
card("NAXIS1 = 8"),
|
||||
card("NAXIS2 = 1"),
|
||||
card("PCOUNT = 0"),
|
||||
card("GCOUNT = 1"),
|
||||
card("TFIELDS = -1"),
|
||||
card("END"),
|
||||
)
|
||||
path := writeHostile(t, "tfields.fits", append(hdr, make([]byte, 2880)...))
|
||||
_, err := LoadFITSTable(path)
|
||||
if err == nil {
|
||||
t.Fatal("LoadFITSTable accepted TFIELDS = -1")
|
||||
}
|
||||
if !strings.Contains(err.Error(), "negative") {
|
||||
t.Fatalf("error = %v, want the negative-count refusal", err)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// TestLoadFITSNaxisRefKeyword pins the NAXIS prefix: any keyword that
|
||||
// started with NAXIS used to count as an axis, so a user keyword
|
||||
// NAXISREF = 7 answered "NAXIS = 1 with 2 NAXISn cards" and refused a
|
||||
// legal file. Only a number after the prefix is an axis, the rule the
|
||||
// writer applies; NAXIS1 keeps counting.
|
||||
func TestLoadFITSNaxisRefKeyword(t *testing.T) {
|
||||
hdr := cardBlock(
|
||||
card("SIMPLE = T"),
|
||||
card("BITPIX = -64"),
|
||||
card("NAXIS = 1"),
|
||||
card("NAXIS1 = 2"),
|
||||
card("NAXISREF= 7"),
|
||||
card("END"),
|
||||
)
|
||||
payload := binary.BigEndian.AppendUint64(nil, math.Float64bits(1.5))
|
||||
payload = binary.BigEndian.AppendUint64(payload, math.Float64bits(-0.5))
|
||||
a, headers, err := LoadFITS(writeHostile(t, "naxisref.fits", append(hdr, payload...)))
|
||||
if err != nil {
|
||||
t.Fatalf("LoadFITS refused a file with a NAXISREF keyword: %v", err)
|
||||
}
|
||||
if a.Len() != 2 || a.FloatAt(0) != 1.5 || a.FloatAt(1) != -0.5 {
|
||||
t.Fatalf("values = %v, want 1.5 and -0.5", a)
|
||||
}
|
||||
if got := headers["NAXISREF"]; got != "7" {
|
||||
t.Fatalf("NAXISREF = %q, want it reported as a user keyword", got)
|
||||
}
|
||||
}
|
||||
+748
@@ -0,0 +1,748 @@
|
||||
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
|
||||
// SPDX-License-Identifier: MIT
|
||||
|
||||
package io
|
||||
|
||||
import (
|
||||
"encoding/binary"
|
||||
"math"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"slices"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"sourcedock.dev/petrbalvin/tensor/internal/core"
|
||||
)
|
||||
|
||||
// The HDF5 fixtures under testdata/h5 were written by the HDF5
|
||||
// reference library, and every value below was read back from them
|
||||
// independently: the expected values in these tests are what the
|
||||
// reference reports, not what this reader
|
||||
// produces.
|
||||
//
|
||||
// fixture.h5: an int32 dataset stored contiguously, a float64 dataset
|
||||
// chunked, gzip compressed and shuffled, a float32 dataset
|
||||
// in a group, and string attributes on the root and the
|
||||
// group (variable-length, so they live in a global heap)
|
||||
// fletcher.h5: a float64 dataset chunked, gzip compressed, with the
|
||||
// fletcher32 checksum filter on top
|
||||
// latest.h5: written with libver="latest", so superblock version 3
|
||||
// and object header version 2, with a float64 dataset /d
|
||||
// and one /g/e in a group
|
||||
|
||||
func h5Fixture(t *testing.T, name string) string {
|
||||
t.Helper()
|
||||
return filepath.Join("testdata", "h5", name)
|
||||
}
|
||||
|
||||
// TestLoadHDF5Values pins the reader against the reference-written fixture.
|
||||
func TestLoadHDF5Values(t *testing.T) {
|
||||
sets, err := LoadHDF5(h5Fixture(t, "fixture.h5"))
|
||||
if err != nil {
|
||||
t.Fatalf("LoadHDF5: %v", err)
|
||||
}
|
||||
if len(sets) != 3 {
|
||||
t.Fatalf("datasets = %d, want 3", len(sets))
|
||||
}
|
||||
byPath := map[string]HDF5Dataset{}
|
||||
for _, d := range sets {
|
||||
byPath[d.Path] = d
|
||||
}
|
||||
// The paths come back sorted.
|
||||
if paths := []string{sets[0].Path, sets[1].Path, sets[2].Path}; paths[0] != "/floats" || paths[1] != "/g/f32" || paths[2] != "/ints" {
|
||||
t.Fatalf("paths = %v, want [/floats /g/f32 /ints]", paths)
|
||||
}
|
||||
|
||||
ints, ok := byPath["/ints"]
|
||||
if !ok {
|
||||
t.Fatal("/ints is missing")
|
||||
}
|
||||
if s := ints.Shape; len(s) != 2 || s[0] != 2 || s[1] != 3 {
|
||||
t.Fatalf("/ints shape = %v, want [2 3]", s)
|
||||
}
|
||||
// The fixture's int32 dataset lands the native int32 dtype: the
|
||||
// reader keeps the width the file stores instead of widening it.
|
||||
if ints.Values.Dtype() != core.Int32 {
|
||||
t.Fatalf("/ints dtype = %s, want int32", ints.Values.Dtype())
|
||||
}
|
||||
for i, want := range []int32{1, 2, 3, 4, 5, 6} {
|
||||
if got := ints.Values.RawInt32s()[i]; got != want {
|
||||
t.Fatalf("/ints[%d] = %d, want %d", i, got, want)
|
||||
}
|
||||
}
|
||||
|
||||
floats, ok := byPath["/floats"]
|
||||
if !ok {
|
||||
t.Fatal("/floats is missing")
|
||||
}
|
||||
if s := floats.Shape; len(s) != 1 || s[0] != 4 {
|
||||
t.Fatalf("/floats shape = %v, want [4]", s)
|
||||
}
|
||||
if floats.Values.Dtype() != core.Float {
|
||||
t.Fatalf("/floats dtype = %s, want float64", floats.Values.Dtype())
|
||||
}
|
||||
for i, want := range []float64{1.5, 2.5, 3.5, 4.5} {
|
||||
if got := floats.Values.RawFloats()[i]; got != want {
|
||||
t.Fatalf("/floats[%d] = %v, want %v", i, got, want)
|
||||
}
|
||||
}
|
||||
|
||||
f32, ok := byPath["/g/f32"]
|
||||
if !ok {
|
||||
t.Fatal("/g/f32 is missing")
|
||||
}
|
||||
if s := f32.Shape; len(s) != 2 || s[0] != 2 || s[1] != 2 {
|
||||
t.Fatalf("/g/f32 shape = %v, want [2 2]", s)
|
||||
}
|
||||
if f32.Values.Dtype() != core.Float32 {
|
||||
t.Fatalf("/g/f32 dtype = %s, want float32", f32.Values.Dtype())
|
||||
}
|
||||
for i, want := range []float32{1, 2, 3, 4} {
|
||||
if got := f32.Values.RawFloat32s()[i]; got != want {
|
||||
t.Fatalf("/g/f32[%d] = %v, want %v", i, got, want)
|
||||
}
|
||||
}
|
||||
|
||||
// The attributes: the root's title reaches every dataset, and the
|
||||
// group's units reach the dataset inside it, the nearest group
|
||||
// winning.
|
||||
if got := ints.Attrs["title"]; got != "h5 fixture" {
|
||||
t.Fatalf("/ints title = %q, want %q", got, "h5 fixture")
|
||||
}
|
||||
if _, ok := ints.Attrs["units"]; ok {
|
||||
t.Fatalf("/ints picked up a group attribute it should not have: %v", ints.Attrs)
|
||||
}
|
||||
if got := f32.Attrs["units"]; got != "K" {
|
||||
t.Fatalf("/g/f32 units = %q, want K", got)
|
||||
}
|
||||
if got := f32.Attrs["title"]; got != "h5 fixture" {
|
||||
t.Fatalf("/g/f32 title = %q, want the inherited one", got)
|
||||
}
|
||||
}
|
||||
|
||||
// TestLoadHDF5Fletcher32 pins the checksum filter: the chunk carries a
|
||||
// fletcher32 sum that must verify before the chunk is used.
|
||||
func TestLoadHDF5Fletcher32(t *testing.T) {
|
||||
sets, err := LoadHDF5(h5Fixture(t, "fletcher.h5"))
|
||||
if err != nil {
|
||||
t.Fatalf("LoadHDF5: %v", err)
|
||||
}
|
||||
if len(sets) != 1 {
|
||||
t.Fatalf("datasets = %d, want 1", len(sets))
|
||||
}
|
||||
d := sets[0]
|
||||
if s := d.Shape; len(s) != 1 || s[0] != 20 {
|
||||
t.Fatalf("shape = %v, want [20]", s)
|
||||
}
|
||||
for i := range 20 {
|
||||
if got := d.Values.RawFloats()[i]; got != float64(i) {
|
||||
t.Fatalf("value %d = %v, want %d", i, got, i)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// TestLoadHDF5Latest pins the "latest" file format against the
|
||||
// reference-written fixture: superblock version 3, object headers
|
||||
// version 2 with their lookup3 checksums, compact groups carrying link
|
||||
// messages, and contiguous datasets.
|
||||
func TestLoadHDF5Latest(t *testing.T) {
|
||||
sets, err := LoadHDF5(h5Fixture(t, "latest.h5"))
|
||||
if err != nil {
|
||||
t.Fatalf("LoadHDF5: %v", err)
|
||||
}
|
||||
if len(sets) != 2 {
|
||||
t.Fatalf("datasets = %d, want 2", len(sets))
|
||||
}
|
||||
d, e := sets[0], sets[1]
|
||||
if d.Path != "/d" || e.Path != "/g/e" {
|
||||
t.Fatalf("paths = %q, %q, want /d and /g/e", d.Path, e.Path)
|
||||
}
|
||||
if s := d.Shape; len(s) != 1 || s[0] != 3 {
|
||||
t.Fatalf("/d shape = %v, want [3]", s)
|
||||
}
|
||||
for i, want := range []float64{1, 2, 3} {
|
||||
if got := d.Values.RawFloats()[i]; got != want {
|
||||
t.Fatalf("/d[%d] = %v, want %v", i, got, want)
|
||||
}
|
||||
}
|
||||
if s := e.Shape; len(s) != 1 || s[0] != 1 {
|
||||
t.Fatalf("/g/e shape = %v, want [1]", s)
|
||||
}
|
||||
if got := e.Values.RawFloats()[0]; got != 4 {
|
||||
t.Fatalf("/g/e[0] = %v, want 4", got)
|
||||
}
|
||||
}
|
||||
|
||||
// TestHDF5Lookup3 pins the checksum against the sums the reference
|
||||
// library wrote into the latest fixture: the superblock's and two
|
||||
// object headers'. The literals are what the file stores, not what
|
||||
// this implementation computes.
|
||||
func TestHDF5Lookup3(t *testing.T) {
|
||||
raw, err := os.ReadFile(h5Fixture(t, "latest.h5"))
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
for _, c := range []struct {
|
||||
name string
|
||||
want uint32
|
||||
lo int
|
||||
hi int
|
||||
}{
|
||||
{"superblock", 0x39ff1913, 0, 44},
|
||||
{"root header", 0xb91c2db3, 48, 175},
|
||||
{"dataset header", 0x8d125cb5, 179, 443},
|
||||
} {
|
||||
if got := hdf5Lookup3(raw[c.lo:c.hi]); got != c.want {
|
||||
t.Errorf("%s: lookup3 = %#08x, want %#08x", c.name, got, c.want)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// TestLoadHDF5Refusals pins the errors: a file that is not HDF5 at
|
||||
// all, a truncated file, and corrupted latest-format checksums must
|
||||
// each be refused with a message that says so, never read halfway.
|
||||
func TestLoadHDF5Refusals(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
notHDF5 := filepath.Join(dir, "plain.bin")
|
||||
if err := os.WriteFile(notHDF5, []byte("this is not an HDF5 file at all, not even close"), 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if _, err := LoadHDF5(notHDF5); err == nil {
|
||||
t.Fatal("expected an error for a file without the HDF5 signature")
|
||||
} else if !strings.Contains(err.Error(), "signature") {
|
||||
t.Fatalf("error = %v, want a signature refusal", err)
|
||||
}
|
||||
// A truncated copy of a valid file: the reader may accept it only
|
||||
// when the structures it actually reads are complete, and it must
|
||||
// never hand back a partial array. Whatever the cut, the call either
|
||||
// errors or returns datasets whose element count matches their
|
||||
// shape.
|
||||
whole, err := os.ReadFile(h5Fixture(t, "fixture.h5"))
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
for _, cut := range []int{8, 32, 100, 600, len(whole) / 2, len(whole) - 4} {
|
||||
path := filepath.Join(dir, "cut.h5")
|
||||
if err := os.WriteFile(path, whole[:cut], 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
sets, err := LoadHDF5(path)
|
||||
if err != nil {
|
||||
continue // refused, which is the expected answer
|
||||
}
|
||||
for _, d := range sets {
|
||||
n := 1
|
||||
for _, s := range d.Shape {
|
||||
n *= s
|
||||
}
|
||||
if d.Values.Len() != n {
|
||||
t.Fatalf("a file truncated to %d bytes gave %q %d values for shape %v",
|
||||
cut, d.Path, d.Values.Len(), d.Shape)
|
||||
}
|
||||
}
|
||||
}
|
||||
// The header itself must be refused: a superblock shorter than its
|
||||
// fixed part cannot be read at all.
|
||||
if _, err := LoadHDF5(writeCut(t, dir, whole, 40)); err == nil {
|
||||
t.Fatal("expected an error for a file truncated inside the superblock")
|
||||
}
|
||||
// The latest format verifies its checksums: a flipped byte in the
|
||||
// superblock and one in an object header must each refuse the file
|
||||
// instead of reading past the corruption.
|
||||
latest, err := os.ReadFile(h5Fixture(t, "latest.h5"))
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
for _, c := range []struct {
|
||||
name string
|
||||
at int
|
||||
}{
|
||||
// Byte 11 is the superblock's consistency flags, which the
|
||||
// reader would otherwise ignore: only the checksum sees it.
|
||||
{"superblock", 11},
|
||||
{"object header", 60},
|
||||
} {
|
||||
corrupt := slices.Clone(latest)
|
||||
corrupt[c.at] ^= 0xff
|
||||
path := filepath.Join(dir, "corrupt.h5")
|
||||
if err := os.WriteFile(path, corrupt, 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if _, err := LoadHDF5(path); err == nil {
|
||||
t.Fatalf("expected an error for a corrupted %s", c.name)
|
||||
} else if !strings.Contains(err.Error(), "checksum") {
|
||||
t.Fatalf("corrupted %s: error = %v, want a checksum refusal", c.name, err)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// writeCut writes the first n bytes of data to a temp file and returns
|
||||
// its path.
|
||||
func writeCut(t *testing.T, dir string, data []byte, n int) string {
|
||||
t.Helper()
|
||||
path := filepath.Join(dir, "cut40.h5")
|
||||
if err := os.WriteFile(path, data[:n], 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
return path
|
||||
}
|
||||
|
||||
// h5FixedType renders a version 1 fixed-point datatype message of the
|
||||
// given element width and signedness. The message carries the bit
|
||||
// offset and bit precision the HDF5 file format specification's
|
||||
// fixed-point property table defines behind the eight-byte header,
|
||||
// twelve bytes in total; the reader keys the landing on the header's
|
||||
// size and signed bit.
|
||||
func h5FixedType(size uint32, signed bool) []byte {
|
||||
m := make([]byte, 12)
|
||||
m[0] = 0x10 // version 1, class 0 (fixed-point)
|
||||
if signed {
|
||||
m[1] = 0x08 // class bit field: bit 3 marks two's complement
|
||||
}
|
||||
binary.LittleEndian.PutUint32(m[4:], size)
|
||||
binary.LittleEndian.PutUint16(m[8:], 0) // bit offset
|
||||
binary.LittleEndian.PutUint16(m[10:], uint16(8*size)) // bit precision
|
||||
return m
|
||||
}
|
||||
|
||||
// h5EnumBoolType renders the boolean enumeration datatype message HDF5
|
||||
// writers carry booleans in, following the HDF5 file format
|
||||
// specification's enumeration class layout: the member count in the
|
||||
// class bit field, the base type as a complete fixed-point message,
|
||||
// each member name NUL-terminated and padded from its own field start
|
||||
// to a multiple of eight bytes, and the packed member values behind
|
||||
// the names.
|
||||
func h5EnumBoolType(names []string, values []byte) []byte {
|
||||
// Version 1, class 8; member count; size 1; then the base type.
|
||||
m := []byte{0x18, byte(len(names)), 0, 0, 1, 0, 0, 0}
|
||||
m = append(m, h5FixedType(1, false)...)
|
||||
for _, n := range names {
|
||||
start := len(m)
|
||||
m = append(m, n...)
|
||||
m = append(m, 0)
|
||||
for (len(m)-start)%8 != 0 {
|
||||
m = append(m, 0)
|
||||
}
|
||||
}
|
||||
m = append(m, values...)
|
||||
return m
|
||||
}
|
||||
|
||||
// h5AttrMessage renders a version 1 attribute message: the name and
|
||||
// every field boundary padded to the eight-byte grid the message
|
||||
// format defines, then the value bytes.
|
||||
func h5AttrMessage(name string, dtypeMsg []byte, dims []uint64, value []byte) []byte {
|
||||
space := h5Dataspace(dims...)
|
||||
nameSize := len(name) + 1
|
||||
dtypeAt := alignUp(8+nameSize, 8)
|
||||
spaceAt := alignUp(dtypeAt+len(dtypeMsg), 8)
|
||||
b := make([]byte, spaceAt+len(space)+len(value))
|
||||
b[0] = 1
|
||||
binary.LittleEndian.PutUint16(b[2:], uint16(nameSize))
|
||||
binary.LittleEndian.PutUint16(b[4:], uint16(len(dtypeMsg)))
|
||||
binary.LittleEndian.PutUint16(b[6:], uint16(len(space)))
|
||||
copy(b[8:], name) // the trailing NUL is the buffer's own zero
|
||||
copy(b[dtypeAt:], dtypeMsg)
|
||||
copy(b[spaceAt:], space)
|
||||
copy(b[spaceAt+len(space):], value)
|
||||
return b
|
||||
}
|
||||
|
||||
// h5ChunkTreeWidth writes a one-entry leaf chunk B-tree for a dataset
|
||||
// of the given rank whose chunk elements are width bytes wide: the
|
||||
// key's element-size slot must agree with the datatype, which the
|
||||
// reader checks.
|
||||
func h5ChunkTreeWidth(f []byte, off, rank int, width uint64, size uint32, chunkAt uint64) {
|
||||
copy(f[off:], hdf5Tree)
|
||||
f[off+4] = 1 // chunk tree
|
||||
f[off+5] = 0 // leaf level
|
||||
binary.LittleEndian.PutUint16(f[off+6:], 1)
|
||||
p := off + 24
|
||||
binary.LittleEndian.PutUint32(f[p:], size)
|
||||
// The filter mask stays zero; the chunk offsets stay zero.
|
||||
binary.LittleEndian.PutUint64(f[p+8+8*rank:], width)
|
||||
binary.LittleEndian.PutUint64(f[p+8+8*(rank+1):], chunkAt)
|
||||
}
|
||||
|
||||
// TestLoadHDF5NativeFixedPoint pins the fixed-point landings of the
|
||||
// contiguous path: every stored width and signedness lands the core
|
||||
// dtype that holds it exactly, extremes included, and int64 stays int.
|
||||
func TestLoadHDF5NativeFixedPoint(t *testing.T) {
|
||||
cases := []struct {
|
||||
name string
|
||||
dtype []byte
|
||||
payload []byte
|
||||
want core.Dtype
|
||||
check func(t *testing.T, a *core.Array)
|
||||
}{
|
||||
{"int8", h5FixedType(1, true), []byte{0x80, 0x00, 0x7f}, core.Int8,
|
||||
func(t *testing.T, a *core.Array) {
|
||||
if got, want := a.RawInt8s()[:3], []int8{-128, 0, 127}; !slices.Equal(got, want) {
|
||||
t.Fatalf("values = %v, want %v", got, want)
|
||||
}
|
||||
}},
|
||||
{"uint8", h5FixedType(1, false), []byte{0x00, 0x01, 0xff}, core.Uint8,
|
||||
func(t *testing.T, a *core.Array) {
|
||||
if got, want := a.RawUint8s()[:3], []uint8{0, 1, 255}; !slices.Equal(got, want) {
|
||||
t.Fatalf("values = %v, want %v", got, want)
|
||||
}
|
||||
}},
|
||||
{"int16", h5FixedType(2, true), []byte{0x00, 0x80, 0xff, 0xff, 0xff, 0x7f}, core.Int16,
|
||||
func(t *testing.T, a *core.Array) {
|
||||
if got, want := a.RawInt16s()[:3], []int16{-32768, -1, 32767}; !slices.Equal(got, want) {
|
||||
t.Fatalf("values = %v, want %v", got, want)
|
||||
}
|
||||
}},
|
||||
{"uint16", h5FixedType(2, false), []byte{0x00, 0x00, 0x00, 0x10, 0xff, 0xff}, core.Uint16,
|
||||
func(t *testing.T, a *core.Array) {
|
||||
if got, want := a.RawUint16s()[:3], []uint16{0, 4096, 65535}; !slices.Equal(got, want) {
|
||||
t.Fatalf("values = %v, want %v", got, want)
|
||||
}
|
||||
}},
|
||||
{"int32", h5FixedType(4, true),
|
||||
[]byte{0x00, 0x00, 0x00, 0x80, 0xff, 0xff, 0xff, 0xff, 0xff, 0xff, 0xff, 0x7f}, core.Int32,
|
||||
func(t *testing.T, a *core.Array) {
|
||||
if got, want := a.RawInt32s()[:3], []int32{-2147483648, -1, 2147483647}; !slices.Equal(got, want) {
|
||||
t.Fatalf("values = %v, want %v", got, want)
|
||||
}
|
||||
}},
|
||||
{"uint32", h5FixedType(4, false),
|
||||
[]byte{0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x40, 0xff, 0xff, 0xff, 0xff}, core.Uint32,
|
||||
func(t *testing.T, a *core.Array) {
|
||||
if got, want := a.RawUint32s()[:3], []uint32{0, 1 << 30, 4294967295}; !slices.Equal(got, want) {
|
||||
t.Fatalf("values = %v, want %v", got, want)
|
||||
}
|
||||
}},
|
||||
{"int64 stays int", h5FixedType(8, true),
|
||||
[]byte{0xfb, 0xff, 0xff, 0xff, 0xff, 0xff, 0xff, 0xff,
|
||||
0, 0, 0, 0, 0, 0, 0, 0,
|
||||
0, 0, 0, 0, 0, 0x01, 0, 0}, core.Int,
|
||||
func(t *testing.T, a *core.Array) {
|
||||
if got, want := a.RawInts()[:3], []int64{-5, 0, 1 << 40}; !slices.Equal(got, want) {
|
||||
t.Fatalf("values = %v, want %v", got, want)
|
||||
}
|
||||
}},
|
||||
}
|
||||
for _, tc := range cases {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
f := h5HostileFile(512)
|
||||
h5ObjectHeader(f, 96,
|
||||
h5Msg{hdf5MsgDataspace, h5Dataspace(3)},
|
||||
h5Msg{hdf5MsgDatatype, tc.dtype},
|
||||
h5Msg{hdf5MsgDataLayout, h5ContiguousLayout(448, uint64(len(tc.payload)))},
|
||||
)
|
||||
copy(f[448:], tc.payload)
|
||||
sets, err := LoadHDF5(writeHostile(t, "native.h5", f))
|
||||
if err != nil {
|
||||
t.Fatalf("LoadHDF5: %v", err)
|
||||
}
|
||||
if len(sets) != 1 {
|
||||
t.Fatalf("datasets = %d, want 1", len(sets))
|
||||
}
|
||||
d := sets[0]
|
||||
if d.Values.Dtype() != tc.want {
|
||||
t.Fatalf("dtype = %s, want %s", d.Values.Dtype(), tc.want)
|
||||
}
|
||||
if s := d.Shape; len(s) != 1 || s[0] != 3 {
|
||||
t.Fatalf("shape = %v, want [3]", s)
|
||||
}
|
||||
tc.check(t, d.Values)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// TestLoadHDF5ChunkedNativeLandings pins the chunked dispatch: the
|
||||
// per-cell decode lands the same native dtypes the contiguous path
|
||||
// lands, through the chunk B-tree and the placement walk.
|
||||
func TestLoadHDF5ChunkedNativeLandings(t *testing.T) {
|
||||
cases := []struct {
|
||||
name string
|
||||
dtype []byte
|
||||
width uint64
|
||||
payload []byte
|
||||
want core.Dtype
|
||||
check func(t *testing.T, a *core.Array)
|
||||
}{
|
||||
{"uint8", h5FixedType(1, false), 1, []byte{0, 1, 255, 42}, core.Uint8,
|
||||
func(t *testing.T, a *core.Array) {
|
||||
if got, want := a.RawUint8s()[:4], []uint8{0, 1, 255, 42}; !slices.Equal(got, want) {
|
||||
t.Fatalf("values = %v, want %v", got, want)
|
||||
}
|
||||
}},
|
||||
{"int16", h5FixedType(2, true), 2,
|
||||
[]byte{0xfd, 0xff, 0x00, 0x80, 0xff, 0x7f, 0x07, 0x00}, core.Int16,
|
||||
func(t *testing.T, a *core.Array) {
|
||||
if got, want := a.RawInt16s()[:4], []int16{-3, -32768, 32767, 7}; !slices.Equal(got, want) {
|
||||
t.Fatalf("values = %v, want %v", got, want)
|
||||
}
|
||||
}},
|
||||
{"bool", h5EnumBoolType([]string{"TRUE", "FALSE"}, []byte{1, 0}), 1,
|
||||
[]byte{1, 0, 1, 1}, core.Bool,
|
||||
func(t *testing.T, a *core.Array) {
|
||||
if got, want := a.RawBools()[:4], []bool{true, false, true, true}; !slices.Equal(got, want) {
|
||||
t.Fatalf("values = %v, want %v", got, want)
|
||||
}
|
||||
}},
|
||||
}
|
||||
for _, tc := range cases {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
const btree, chunkAt = 256, 320
|
||||
f := h5HostileFile(512)
|
||||
end := h5ObjectHeader(f, 96,
|
||||
h5Msg{hdf5MsgDataspace, h5Dataspace(4)},
|
||||
h5Msg{hdf5MsgDatatype, tc.dtype},
|
||||
h5Msg{hdf5MsgDataLayout, h5ChunkLayoutV3(btree, 4, uint32(tc.width))},
|
||||
)
|
||||
if end > btree {
|
||||
t.Fatalf("the test object header runs to %d, past the chunk B-tree at %d", end, btree)
|
||||
}
|
||||
h5ChunkTreeWidth(f, btree, 1, tc.width, uint32(len(tc.payload)), chunkAt)
|
||||
copy(f[chunkAt:], tc.payload)
|
||||
sets, err := LoadHDF5(writeHostile(t, "chunk-native.h5", f))
|
||||
if err != nil {
|
||||
t.Fatalf("LoadHDF5: %v", err)
|
||||
}
|
||||
if len(sets) != 1 {
|
||||
t.Fatalf("datasets = %d, want 1", len(sets))
|
||||
}
|
||||
d := sets[0]
|
||||
if d.Values.Dtype() != tc.want {
|
||||
t.Fatalf("dtype = %s, want %s", d.Values.Dtype(), tc.want)
|
||||
}
|
||||
tc.check(t, d.Values)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// enumBoolLoad builds a one-dataset contiguous file around a datatype
|
||||
// message and payload, and returns the load error or the dataset.
|
||||
func enumBoolLoad(t *testing.T, dtypeMsg, payload []byte) ([]HDF5Dataset, error) {
|
||||
t.Helper()
|
||||
f := h5HostileFile(512)
|
||||
h5ObjectHeader(f, 96,
|
||||
h5Msg{hdf5MsgDataspace, h5Dataspace(uint64(len(payload)))},
|
||||
h5Msg{hdf5MsgDatatype, dtypeMsg},
|
||||
h5Msg{hdf5MsgDataLayout, h5ContiguousLayout(448, uint64(len(payload)))},
|
||||
)
|
||||
copy(f[448:], payload)
|
||||
return LoadHDF5(writeHostile(t, "enum.h5", f))
|
||||
}
|
||||
|
||||
// TestLoadHDF5EnumBoolLandings pins the boolean enumeration landing:
|
||||
// a one-byte unsigned base whose member values are a subset of {0, 1}
|
||||
// lands core.Bool whatever the member names say, because the values,
|
||||
// not the names, carry the semantics.
|
||||
func TestLoadHDF5EnumBoolLandings(t *testing.T) {
|
||||
cases := []struct {
|
||||
name string
|
||||
dtype []byte
|
||||
values []byte
|
||||
want []bool
|
||||
}{
|
||||
{"members TRUE and FALSE",
|
||||
h5EnumBoolType([]string{"TRUE", "FALSE"}, []byte{1, 0}),
|
||||
[]byte{1, 0, 1}, []bool{true, false, true}},
|
||||
{"names are irrelevant to the values",
|
||||
h5EnumBoolType([]string{"present", "absent"}, []byte{0, 1}),
|
||||
[]byte{0, 1, 1}, []bool{false, true, true}},
|
||||
{"a single member of zero",
|
||||
h5EnumBoolType([]string{"off"}, []byte{0}),
|
||||
[]byte{0, 0, 0}, []bool{false, false, false}},
|
||||
}
|
||||
for _, tc := range cases {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
sets, err := enumBoolLoad(t, tc.dtype, tc.values)
|
||||
if err != nil {
|
||||
t.Fatalf("LoadHDF5: %v", err)
|
||||
}
|
||||
if len(sets) != 1 {
|
||||
t.Fatalf("datasets = %d, want 1", len(sets))
|
||||
}
|
||||
d := sets[0]
|
||||
if d.Values.Dtype() != core.Bool {
|
||||
t.Fatalf("dtype = %s, want bool", d.Values.Dtype())
|
||||
}
|
||||
if got := d.Values.RawBools()[:len(tc.want)]; !slices.Equal(got, tc.want) {
|
||||
t.Fatalf("values = %v, want %v", got, tc.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// TestLoadHDF5EnumRefusals pins the loud refusals: every enumeration
|
||||
// outside the boolean convention, every bit field, and a boolean
|
||||
// payload cell outside the members are named errors, never a silent
|
||||
// guess. The variants mutate the spec-shaped message, which also pins
|
||||
// the field offsets the parser reads.
|
||||
func TestLoadHDF5EnumRefusals(t *testing.T) {
|
||||
signed := func() []byte { m := h5EnumBoolType([]string{"TRUE", "FALSE"}, []byte{1, 0}); m[9] |= 0x08; return m }
|
||||
bigEndian := func() []byte { m := h5EnumBoolType([]string{"TRUE", "FALSE"}, []byte{1, 0}); m[9] |= 0x01; return m }
|
||||
baseClass := func() []byte { m := h5EnumBoolType([]string{"TRUE", "FALSE"}, []byte{1, 0}); m[8] = 0x11; return m }
|
||||
baseSize := func() []byte {
|
||||
m := h5EnumBoolType([]string{"TRUE", "FALSE"}, []byte{1, 0})
|
||||
binary.LittleEndian.PutUint32(m[12:], 2)
|
||||
return m
|
||||
}
|
||||
valueSize := func() []byte {
|
||||
m := h5EnumBoolType([]string{"TRUE", "FALSE"}, []byte{1, 0})
|
||||
binary.LittleEndian.PutUint32(m[4:], 2)
|
||||
return m
|
||||
}
|
||||
noMembers := func() []byte { m := h5EnumBoolType([]string{"TRUE", "FALSE"}, []byte{1, 0}); m[1] = 0; return m }
|
||||
reserved := func() []byte { m := h5EnumBoolType([]string{"TRUE", "FALSE"}, []byte{1, 0}); m[3] = 0x04; return m }
|
||||
shortNames := func() []byte { m := h5EnumBoolType([]string{"TRUE", "FALSE"}, []byte{1, 0}); m[1] = 3; return m }
|
||||
bitField := func() []byte { m := h5FixedType(1, false); m[0] = 0x14; return m }
|
||||
|
||||
cases := []struct {
|
||||
name string
|
||||
dtype []byte
|
||||
payload []byte
|
||||
want string
|
||||
}{
|
||||
{"a member value outside {0, 1}", h5EnumBoolType([]string{"A", "B"}, []byte{0, 2}), []byte{0, 1}, "outside the boolean convention"},
|
||||
{"a signed base type", signed(), []byte{1, 0}, "signed base type"},
|
||||
{"a big-endian base type", bigEndian(), []byte{1, 0}, "big-endian"},
|
||||
{"a non-fixed-point base type", baseClass(), []byte{1, 0}, "base type of class 1"},
|
||||
{"a base type wider than one byte", baseSize(), []byte{1, 0}, "base type of 2 bytes"},
|
||||
{"values wider than one byte", valueSize(), []byte{1, 0}, "2-byte values"},
|
||||
{"no members at all", noMembers(), []byte{0}, "declares 0 members"},
|
||||
{"reserved bit field bits", reserved(), []byte{1, 0}, "unknown bit field bits"},
|
||||
{"more members than names", shortNames(), []byte{1, 0}, "ends inside"},
|
||||
{"a bit field datatype", bitField(), []byte{0}, "bit field"},
|
||||
{"a payload cell outside the members",
|
||||
h5EnumBoolType([]string{"TRUE", "FALSE"}, []byte{1, 0}), []byte{1, 7, 0}, "outside the members 0 and 1"},
|
||||
}
|
||||
for _, tc := range cases {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
sets, err := enumBoolLoad(t, tc.dtype, tc.payload)
|
||||
if err == nil {
|
||||
t.Fatalf("LoadHDF5 accepted %s: %d datasets, %v", tc.name, len(sets), sets)
|
||||
}
|
||||
if !strings.Contains(err.Error(), tc.want) {
|
||||
t.Fatalf("error = %v, want it to carry %q", err, tc.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
// The same refusal on the chunked path, where the per-cell decode
|
||||
// runs inside the placement walk.
|
||||
t.Run("a chunked payload cell outside the members", func(t *testing.T) {
|
||||
const btree, chunkAt = 256, 320
|
||||
f := h5HostileFile(512)
|
||||
end := h5ObjectHeader(f, 96,
|
||||
h5Msg{hdf5MsgDataspace, h5Dataspace(4)},
|
||||
h5Msg{hdf5MsgDatatype, h5EnumBoolType([]string{"TRUE", "FALSE"}, []byte{1, 0})},
|
||||
h5Msg{hdf5MsgDataLayout, h5ChunkLayoutV3(btree, 4, 1)},
|
||||
)
|
||||
if end > btree {
|
||||
t.Fatalf("the test object header runs to %d, past the chunk B-tree at %d", end, btree)
|
||||
}
|
||||
h5ChunkTreeWidth(f, btree, 1, 1, 4, chunkAt)
|
||||
copy(f[chunkAt:], []byte{1, 9, 0, 1})
|
||||
_, err := LoadHDF5(writeHostile(t, "enum-chunk.h5", f))
|
||||
if err == nil || !strings.Contains(err.Error(), "outside the members 0 and 1") {
|
||||
t.Fatalf("chunked enum payload of 9: err = %v, want the member refusal", err)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
// TestLoadHDF5Uint64Refused pins the unsigned 64-bit refusal at both
|
||||
// sites that hold it: the dataset gate and the value decode, with the
|
||||
// same text at each.
|
||||
func TestLoadHDF5Uint64Refused(t *testing.T) {
|
||||
const want = "unsigned 64-bit integers have no exact core dtype"
|
||||
f := h5HostileFile(512)
|
||||
h5ObjectHeader(f, 96,
|
||||
h5Msg{hdf5MsgDataspace, h5Dataspace(1)},
|
||||
h5Msg{hdf5MsgDatatype, h5FixedType(8, false)},
|
||||
h5Msg{hdf5MsgDataLayout, h5ContiguousLayout(448, 8)},
|
||||
)
|
||||
copy(f[448:], []byte{1, 0, 0, 0, 0, 0, 0, 0})
|
||||
_, err := LoadHDF5(writeHostile(t, "uint64.h5", f))
|
||||
if err == nil || !strings.Contains(err.Error(), want) {
|
||||
t.Fatalf("LoadHDF5 on an unsigned 64-bit dataset: err = %v, want it to carry %q", err, want)
|
||||
}
|
||||
// The decode site, called directly: the same text, no widening.
|
||||
if _, err := arrayFromRaw(make([]byte, 8), hdf5Type{class: 0, size: 8, width: 8}, []int{1}); err == nil ||
|
||||
!strings.Contains(err.Error(), want) {
|
||||
t.Fatalf("arrayFromRaw on unsigned 64-bit bytes: err = %v, want it to carry %q", err, want)
|
||||
}
|
||||
// The chunked dispatch, through a dataset fixture: the chunk
|
||||
// dimensions carry the element size in their last slot, matching
|
||||
// the datatype, and the chunk B-tree address points past the end
|
||||
// of the file, so the pin also records that the dataset gate
|
||||
// refuses the datatype before any storage or tree is read.
|
||||
cf := h5HostileFile(512)
|
||||
h5ObjectHeader(cf, 96,
|
||||
h5Msg{hdf5MsgDataspace, h5Dataspace(1)},
|
||||
h5Msg{hdf5MsgDatatype, h5FixedType(8, false)},
|
||||
h5Msg{hdf5MsgDataLayout, h5ChunkLayoutV3(1024, 1, 8)},
|
||||
)
|
||||
if _, err := LoadHDF5(writeHostile(t, "uint64-chunked.h5", cf)); err == nil ||
|
||||
!strings.Contains(err.Error(), want) {
|
||||
t.Fatalf("LoadHDF5 on a chunked unsigned 64-bit dataset: err = %v, want it to carry %q", err, want)
|
||||
}
|
||||
// The chunkedArray dispatch itself, called directly with width 8
|
||||
// unsigned: the same refusal, reached before any chunk walk.
|
||||
var fh hdf5File
|
||||
if _, err := fh.chunkedArray("/u64", []int{1},
|
||||
hdf5Type{class: 0, size: 8, width: 8},
|
||||
hdf5Layout{class: 2, dims: []int{1}}, nil, 8); err == nil ||
|
||||
!strings.Contains(err.Error(), want) {
|
||||
t.Fatalf("chunkedArray on unsigned 64-bit bytes: err = %v, want it to carry %q", err, want)
|
||||
}
|
||||
}
|
||||
|
||||
// TestLoadHDF5AttributeSignedRendering pins the attribute text of
|
||||
// numeric attributes: the datatype's own signed bit decides how its
|
||||
// stored bits read, the boolean enumeration renders 0 and 1, unsigned
|
||||
// values keep every digit, and a cell outside the boolean members
|
||||
// drops the attribute instead of guessing.
|
||||
func TestLoadHDF5AttributeSignedRendering(t *testing.T) {
|
||||
const datasetAt, dataAt = 768, 832
|
||||
f := h5HostileFile(896)
|
||||
msgs := []h5Msg{{hdf5MsgLink, h5HardLink("d", datasetAt)}}
|
||||
// The dataspace carries the element count; the value bytes follow
|
||||
// it as count many datatype-width cells.
|
||||
add := func(name string, dtypeMsg, value []byte, elems uint64) {
|
||||
msgs = append(msgs, h5Msg{hdf5MsgAttribute, h5AttrMessage(name, dtypeMsg, []uint64{elems}, value)})
|
||||
}
|
||||
add("s8", h5FixedType(1, true), []byte{0xff}, 1)
|
||||
add("u8", h5FixedType(1, false), []byte{0xff}, 1)
|
||||
add("s16", h5FixedType(2, true), []byte{0xfe, 0xff}, 1)
|
||||
add("u32", h5FixedType(4, false), []byte{0xff, 0xff, 0xff, 0xff}, 1)
|
||||
add("flag", h5EnumBoolType([]string{"TRUE", "FALSE"}, []byte{1, 0}), []byte{1, 0}, 2)
|
||||
add("bad", h5EnumBoolType([]string{"TRUE", "FALSE"}, []byte{1, 0}), []byte{1, 7}, 2)
|
||||
end := h5ObjectHeader(f, 96, msgs...)
|
||||
if end > datasetAt {
|
||||
t.Fatalf("the root header runs to %d, past the dataset at %d", end, datasetAt)
|
||||
}
|
||||
h5ObjectHeader(f, datasetAt,
|
||||
h5Msg{hdf5MsgDataspace, h5Dataspace(1)},
|
||||
h5Msg{hdf5MsgDatatype, h5FloatType(8)},
|
||||
h5Msg{hdf5MsgDataLayout, h5ContiguousLayout(dataAt, 8)},
|
||||
)
|
||||
binary.LittleEndian.PutUint64(f[dataAt:], math.Float64bits(2.5))
|
||||
sets, err := LoadHDF5(writeHostile(t, "attrs-signed.h5", f))
|
||||
if err != nil {
|
||||
t.Fatalf("LoadHDF5: %v", err)
|
||||
}
|
||||
if len(sets) != 1 || sets[0].Path != "/d" {
|
||||
t.Fatalf("datasets = %v, want the linked /d", sets)
|
||||
}
|
||||
attrs := sets[0].Attrs
|
||||
for k, want := range map[string]string{
|
||||
"s8": "-1", "u8": "255", "s16": "-2", "u32": "4294967295", "flag": "[1, 0]",
|
||||
} {
|
||||
if got := attrs[k]; got != want {
|
||||
t.Fatalf("attr %s = %q, want %q", k, got, want)
|
||||
}
|
||||
}
|
||||
if v, ok := attrs["bad"]; ok {
|
||||
t.Fatalf("the attribute with a cell outside the members was accepted as %q", v)
|
||||
}
|
||||
if got := sets[0].Values.FloatAt(0); got != 2.5 {
|
||||
t.Fatalf("value = %v, want 2.5", got)
|
||||
}
|
||||
}
|
||||
+786
@@ -0,0 +1,786 @@
|
||||
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
|
||||
// SPDX-License-Identifier: MIT
|
||||
|
||||
package io
|
||||
|
||||
// HDF5 writing, the mirror of the reader in hdf5.go. The reader's
|
||||
// verified decoders are the specification: every structure here is
|
||||
// written in the shape the reader accepts and in the shape the
|
||||
// reference library writes, as the fixtures under testdata/h5 pin it.
|
||||
//
|
||||
// Written: superblock version 0 (the classic layout, the default) and
|
||||
// version 3 (the "latest" layout, whose superblock and object headers
|
||||
// carry lookup3 checksums), object headers versions 1 and 2, groups
|
||||
// stored as symbol tables (local heap, version 1 group B-tree, symbol
|
||||
// table nodes) in the classic layout or as link messages in the latest
|
||||
// one, datasets stored contiguously or in chunks through a version 1
|
||||
// chunk B-tree, the deflate and shuffle filters, fixed-point and
|
||||
// floating-point datatypes of the usual widths, the boolean
|
||||
// enumeration convention, fixed-length string datatypes, and
|
||||
// attributes in the object header.
|
||||
//
|
||||
// Every address and length is eight bytes, as in the fixtures. The
|
||||
// output is deterministic: children are written in sorted name order
|
||||
// and nothing depends on map iteration.
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"compress/zlib"
|
||||
"encoding/binary"
|
||||
"maps"
|
||||
"math"
|
||||
"os"
|
||||
"slices"
|
||||
"strings"
|
||||
|
||||
"sourcedock.dev/petrbalvin/tensor/internal/base"
|
||||
"sourcedock.dev/petrbalvin/tensor/internal/core"
|
||||
)
|
||||
|
||||
// HDF5WriteOptions tunes SaveHDF5 and SaveHDF5Text. The zero value
|
||||
// writes the classic file layout with contiguous datasets, which every
|
||||
// reader of the format understands.
|
||||
type HDF5WriteOptions struct {
|
||||
// Latest writes superblock version 3 with version 2 object
|
||||
// headers: groups become link messages and every structure the
|
||||
// format checksums carries a lookup3 sum. Latest files hold
|
||||
// contiguous datasets only: the reference library stores filtered
|
||||
// chunks of a latest file in a version 2 B-tree, which LoadHDF5
|
||||
// does not read, so combining Latest with a filter is refused.
|
||||
Latest bool
|
||||
|
||||
// Gzip applies the deflate filter to every numeric dataset at the
|
||||
// given level: 0 (the default) disables it, -1 means the default
|
||||
// level and 1 to 9 are the levels of the format. A filtered
|
||||
// dataset is stored in chunks.
|
||||
Gzip int
|
||||
|
||||
// Shuffle applies the shuffle filter before deflate, which
|
||||
// regroups the bytes of each element so compression sees the
|
||||
// high-order bytes together. Shuffle alone also forces chunks.
|
||||
Shuffle bool
|
||||
|
||||
// ChunkBytes is the target size of one chunk in bytes for
|
||||
// filtered datasets; 0 selects a default of 64 KiB. Datasets
|
||||
// smaller than the target stay in one chunk.
|
||||
ChunkBytes int
|
||||
}
|
||||
|
||||
// SaveHDF5 writes the datasets as an HDF5 file: the mirror of
|
||||
// LoadHDF5. The paths build the group tree (the dataset "/g/f32" sits
|
||||
// in the group "/g"), so the file reads back with the same paths,
|
||||
// shapes, dtypes and values. Each dataset's Attrs are written on the
|
||||
// dataset itself; the attributes of the root and of the groups come
|
||||
// from groupAttrs, keyed by group path with the root keyed "/".
|
||||
//
|
||||
// Every dtype the writer stores lands the same dtype through LoadHDF5:
|
||||
// bool through the HDF5 boolean enumeration convention, the narrow
|
||||
// integers at their stored width and signedness, float32, float64 and
|
||||
// int64 directly. Float16 and complex arrays are refused: LoadHDF5
|
||||
// decodes neither a two-byte floating-point nor a complex datatype, so
|
||||
// the writer refuses them rather than write a file this package cannot
|
||||
// read back. Attribute values are parsed back into typed attributes: a
|
||||
// whole number becomes an int64 attribute, a decimal a float64 one, a
|
||||
// bracketed list an int64 or float64 array, and anything else a
|
||||
// fixed-length string, so a file written from LoadHDF5's own output
|
||||
// reads back with the same attribute text. When several option values
|
||||
// are passed the last one wins.
|
||||
func SaveHDF5(path string, datasets []HDF5Dataset, groupAttrs map[string]map[string]string, opts ...HDF5WriteOptions) error {
|
||||
const name = "SaveHDF5"
|
||||
options := HDF5WriteOptions{}
|
||||
for _, o := range opts {
|
||||
options = o
|
||||
}
|
||||
if err := hdf5CheckOptions(name, options); err != nil {
|
||||
return err
|
||||
}
|
||||
root, err := hdf5BuildPlan(name, datasets, nil, groupAttrs)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
w := &hdf5Writer{latest: options.Latest, opts: options}
|
||||
if err := w.write(root); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := os.WriteFile(path, w.buf, 0o644); err != nil {
|
||||
return base.Errf("%s: %w", name, err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// HDF5TextDataset is one fixed-length string dataset for SaveHDF5Text:
|
||||
// Text holds the elements in row-major order, each padded to the
|
||||
// longest element in the file. The reader of this package refuses
|
||||
// string datasets (it reads numeric arrays only), so these files are
|
||||
// for other readers of the format.
|
||||
type HDF5TextDataset struct {
|
||||
Path string
|
||||
Shape []int
|
||||
Text []string
|
||||
}
|
||||
|
||||
// SaveHDF5Text writes fixed-length string datasets, the string side of
|
||||
// the datatype message the reader refuses for data but accepts for
|
||||
// attributes. The datasets are stored contiguously; the deflate and
|
||||
// shuffle filters, which are chunked-storage filters, are refused
|
||||
// here by name rather than silently dropped. When several option
|
||||
// values are passed the last one wins.
|
||||
func SaveHDF5Text(path string, texts []HDF5TextDataset, opts ...HDF5WriteOptions) error {
|
||||
const name = "SaveHDF5Text"
|
||||
options := HDF5WriteOptions{}
|
||||
for _, o := range opts {
|
||||
options = o
|
||||
}
|
||||
if options.Gzip != 0 || options.Shuffle {
|
||||
return base.Errf("%s: the deflate and shuffle filters apply to numeric datasets; string datasets are written contiguously", name)
|
||||
}
|
||||
if err := hdf5CheckOptions(name, options); err != nil {
|
||||
return err
|
||||
}
|
||||
root, err := hdf5BuildPlan(name, nil, texts, nil)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
w := &hdf5Writer{latest: options.Latest, opts: options}
|
||||
if err := w.write(root); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := os.WriteFile(path, w.buf, 0o644); err != nil {
|
||||
return base.Errf("%s: %w", name, err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// hdf5CheckOptions rejects the option combinations the writer cannot
|
||||
// honour, naming each one.
|
||||
func hdf5CheckOptions(name string, opts HDF5WriteOptions) error {
|
||||
switch opts.Gzip {
|
||||
case 0, -1, 1, 2, 3, 4, 5, 6, 7, 8, 9:
|
||||
default:
|
||||
return base.Errf("%s: gzip level %d: use 0 to disable the filter, -1 for the default level or 1 to 9", name, opts.Gzip)
|
||||
}
|
||||
if opts.ChunkBytes < 0 {
|
||||
return base.Errf("%s: a chunk target of %d bytes is negative", name, opts.ChunkBytes)
|
||||
}
|
||||
if opts.Latest && (opts.Gzip != 0 || opts.Shuffle) {
|
||||
return base.Errf("%s: the latest format stores filtered chunks through a version 2 B-tree, which LoadHDF5 does not read; write filtered datasets to a classic file", name)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// The B-tree and symbol node fanouts the classic layout declares in
|
||||
// its superblock: a node of the format holds twice the K of its kind,
|
||||
// so a symbol table node holds eight entries, a group B-tree node
|
||||
// thirty-two children and a chunk B-tree node (whose K the format
|
||||
// fixes at thirty-two for superblock version 0) sixty-four chunks.
|
||||
const (
|
||||
hdf5GroupLeafK = 4
|
||||
hdf5GroupInnerK = 16
|
||||
hdf5IStoreK = 32
|
||||
hdf5ChunkTarget = 64 << 10
|
||||
hdf5MaxChunks = 4 << 20
|
||||
hdf5MaxAttrs = 4096
|
||||
hdf5MaxRank = 32
|
||||
hdf5MaxNameBytes = 4096
|
||||
)
|
||||
|
||||
// hdf5OutSet is one dataset of the plan: the source of its payload,
|
||||
// its shape and its datatype class (0 fixed-point, 1 floating-point,
|
||||
// 3 string, 8 the boolean enumeration).
|
||||
type hdf5OutSet struct {
|
||||
path string
|
||||
shape []int
|
||||
class byte
|
||||
width int
|
||||
// signed marks a fixed-point payload as two's complement: the class
|
||||
// bit field's bit 0x08, the bit the datatype message carries and
|
||||
// the reader keys its landing on.
|
||||
signed bool
|
||||
// nbytes is the payload's byte size, which the plan has bounded.
|
||||
nbytes int
|
||||
// The source lives in exactly one of the payload fields below, the
|
||||
// one its class, width and signedness name: a contiguous write
|
||||
// encodes the values straight into the image and a chunked write
|
||||
// stages them chunk by chunk, so no encoded copy of the payload is
|
||||
// built. Every numeric payload is serialised little-endian at its
|
||||
// own stored width, bool as one zero-or-one byte per element.
|
||||
bools []bool
|
||||
ints []int64
|
||||
i8s []int8
|
||||
u8s []uint8
|
||||
i16s []int16
|
||||
u16s []uint16
|
||||
i32s []int32
|
||||
u32s []uint32
|
||||
f32s []float32
|
||||
f64s []float64
|
||||
texts []string
|
||||
attrs []hdf5Attr
|
||||
// written state
|
||||
addr uint64
|
||||
}
|
||||
|
||||
// hdf5OutNode is one group of the plan: the root is the node whose
|
||||
// path is "/". Groups sort their children by name before writing so
|
||||
// the heap offsets and the B-tree order agree.
|
||||
type hdf5OutNode struct {
|
||||
path string
|
||||
name string
|
||||
attrs []hdf5Attr
|
||||
groups []*hdf5OutNode
|
||||
sets []*hdf5OutSet
|
||||
// written state: the object header address and, in the classic
|
||||
// layout, the group's B-tree and local heap.
|
||||
addr uint64
|
||||
btree uint64
|
||||
heap uint64
|
||||
}
|
||||
|
||||
// hdf5BuildPlan validates the paths, the dtypes and the attribute
|
||||
// texts and lays the file out as a tree: the datasets under their
|
||||
// groups, every attribute sorted by name. Anything the writer would
|
||||
// refuse it refuses here, before a byte is written.
|
||||
func hdf5BuildPlan(name string, datasets []HDF5Dataset, texts []HDF5TextDataset, groupAttrs map[string]map[string]string) (*hdf5OutNode, error) {
|
||||
root := &hdf5OutNode{path: "/"}
|
||||
groups := map[string]*hdf5OutNode{"/": root}
|
||||
used := map[string]bool{"/": true}
|
||||
var ensure func(path string) (*hdf5OutNode, error)
|
||||
ensure = func(path string) (*hdf5OutNode, error) {
|
||||
if g, ok := groups[path]; ok {
|
||||
return g, nil
|
||||
}
|
||||
segs, err := hdf5CheckPath(name, path, "group")
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
parentPath := "/"
|
||||
if len(segs) > 1 {
|
||||
parentPath = "/" + strings.Join(segs[:len(segs)-1], "/")
|
||||
}
|
||||
parent, err := ensure(parentPath)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if used[path] {
|
||||
return nil, base.Errf("%s: %q names both a dataset and a group", name, path)
|
||||
}
|
||||
g := &hdf5OutNode{path: path, name: segs[len(segs)-1]}
|
||||
parent.groups = append(parent.groups, g)
|
||||
groups[path] = g
|
||||
used[path] = true
|
||||
return g, nil
|
||||
}
|
||||
place := func(path string) (*hdf5OutNode, error) {
|
||||
segs, err := hdf5CheckPath(name, path, "dataset")
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if used[path] {
|
||||
return nil, base.Errf("%s: the path %q is written twice", name, path)
|
||||
}
|
||||
parentPath := "/"
|
||||
if len(segs) > 1 {
|
||||
parentPath = "/" + strings.Join(segs[:len(segs)-1], "/")
|
||||
}
|
||||
parent, err := ensure(parentPath)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
used[path] = true
|
||||
return parent, nil
|
||||
}
|
||||
for i := range datasets {
|
||||
d := &datasets[i]
|
||||
parent, err := place(d.Path)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
s, err := hdf5PlanSet(name, d)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
parent.sets = append(parent.sets, s)
|
||||
}
|
||||
for i := range texts {
|
||||
tx := &texts[i]
|
||||
parent, err := place(tx.Path)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
s, err := hdf5PlanText(name, tx)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
parent.sets = append(parent.sets, s)
|
||||
}
|
||||
keys := slices.Sorted(maps.Keys(groupAttrs))
|
||||
for _, k := range keys {
|
||||
g, ok := groups[k]
|
||||
if !ok {
|
||||
return nil, base.Errf("%s: the attribute path %q does not name a group of this file", name, k)
|
||||
}
|
||||
attrs, err := hdf5PlanAttrs(name, k, groupAttrs[k])
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
g.attrs = attrs
|
||||
}
|
||||
hdf5SortNode(root)
|
||||
return root, nil
|
||||
}
|
||||
|
||||
// hdf5CheckPath splits an absolute object path into its segments and
|
||||
// refuses what no file should carry: a relative path, the root path
|
||||
// where an object is wanted, an empty segment, a segment holding a
|
||||
// NUL byte or a path longer than the bound a sane file keeps.
|
||||
func hdf5CheckPath(name, path, kind string) ([]string, error) {
|
||||
if path == "" || path[0] != '/' {
|
||||
return nil, base.Errf("%s: %s path %q is not absolute", name, kind, path)
|
||||
}
|
||||
if path == "/" {
|
||||
return nil, base.Errf("%s: the root path does not name a %s", name, kind)
|
||||
}
|
||||
if len(path) > hdf5MaxNameBytes {
|
||||
return nil, base.Errf("%s: %s path %q is longer than %d bytes", name, kind, path, hdf5MaxNameBytes)
|
||||
}
|
||||
segs := strings.Split(path[1:], "/")
|
||||
for _, s := range segs {
|
||||
if s == "" {
|
||||
return nil, base.Errf("%s: %s path %q has an empty segment", name, kind, path)
|
||||
}
|
||||
if strings.IndexByte(s, 0) >= 0 {
|
||||
return nil, base.Errf("%s: %s path %q holds a NUL byte", name, kind, path)
|
||||
}
|
||||
}
|
||||
return segs, nil
|
||||
}
|
||||
|
||||
// hdf5SortNode orders the children of a group by name, recursively, so
|
||||
// the written file is independent of the order the caller supplied.
|
||||
func hdf5SortNode(g *hdf5OutNode) {
|
||||
slices.SortFunc(g.groups, func(a, b *hdf5OutNode) int { return strings.Compare(a.name, b.name) })
|
||||
slices.SortFunc(g.sets, func(a, b *hdf5OutSet) int { return strings.Compare(a.path, b.path) })
|
||||
for _, sub := range g.groups {
|
||||
hdf5SortNode(sub)
|
||||
}
|
||||
}
|
||||
|
||||
// hdf5ImageEstimate bounds the byte size of the image the plan writes,
|
||||
// the number the buffer is allocated from. The bound is loose on
|
||||
// purpose: every structure the format wraps around a payload is
|
||||
// charged a fixed frame, a deflated chunk cannot grow past its own
|
||||
// bytes, and a chunked dataset is charged one chunk of padding plus a
|
||||
// node per sixty-four chunks. Over-estimating costs the memory the
|
||||
// write frees again; under-estimating costs one reallocation.
|
||||
func hdf5ImageEstimate(g *hdf5OutNode, opts HDF5WriteOptions) int {
|
||||
// The sum is kept in int64 and clamped to what an int holds, so no
|
||||
// pile of frames can wrap it into a negative capacity.
|
||||
return int(min(hdf5NodeEstimate(g, opts), maxInt))
|
||||
}
|
||||
|
||||
// maxInt is the largest value an int holds on this platform.
|
||||
const maxInt = int64(^uint(0) >> 1)
|
||||
|
||||
// hdf5NodeEstimate sums one group's own frame, its attributes and the
|
||||
// datasets and subgroups beneath it.
|
||||
func hdf5NodeEstimate(g *hdf5OutNode, opts HDF5WriteOptions) int64 {
|
||||
total := int64(4096+128*(len(g.groups)+len(g.sets))) + hdf5AttrsEstimate(g.attrs)
|
||||
for _, s := range g.sets {
|
||||
total += hdf5SetEstimate(s, opts)
|
||||
}
|
||||
for _, sub := range g.groups {
|
||||
total += hdf5NodeEstimate(sub, opts)
|
||||
}
|
||||
return total
|
||||
}
|
||||
|
||||
// hdf5SetEstimate bounds the image bytes one dataset occupies: its
|
||||
// stored payload, its object header and, when it is chunked, the
|
||||
// padding, the chunk B-tree nodes and the filter that shrinks it.
|
||||
func hdf5SetEstimate(s *hdf5OutSet, opts HDF5WriteOptions) int64 {
|
||||
total := int64(s.nbytes+s.nbytes/64) + 4096
|
||||
if len(s.shape) == 0 || s.nbytes == 0 || s.class == 3 || (opts.Gzip == 0 && !opts.Shuffle) {
|
||||
return total
|
||||
}
|
||||
chunkTarget := hdf5ChunkTargetOf(opts)
|
||||
// The node count is a starting hint the write grows past by append,
|
||||
// never a bound it must honour, so it is capped at what a 512-byte
|
||||
// target would charge: a pathological chunk target of one or two
|
||||
// bytes would otherwise charge a node, and its kilobyte of image,
|
||||
// to every single data byte up front.
|
||||
nodes := 1 + min(s.nbytes/max(chunkTarget/2, 1), s.nbytes/512+1, 1<<20)
|
||||
return total + int64(chunkTarget) + int64(nodes)*int64(hdf5ChunkNodeSize(len(s.shape)))
|
||||
}
|
||||
|
||||
// hdf5AttrsEstimate bounds the attribute messages of one object: every
|
||||
// value's bytes are at most four times the text they are parsed from,
|
||||
// and the fields around them are charged a fixed frame.
|
||||
func hdf5AttrsEstimate(attrs []hdf5Attr) int64 {
|
||||
var total int64
|
||||
for _, a := range attrs {
|
||||
total += int64(4*(len(a.name)+len(a.text)) + 256)
|
||||
}
|
||||
return total
|
||||
}
|
||||
|
||||
// hdf5PlanSet validates one numeric dataset and keeps its values for
|
||||
// the write, which encodes them little-endian, the byte order the
|
||||
// datatype message declares.
|
||||
func hdf5PlanSet(name string, d *HDF5Dataset) (*hdf5OutSet, error) {
|
||||
if d.Values == nil {
|
||||
return nil, base.Errf("%s: dataset %q has no values", name, d.Path)
|
||||
}
|
||||
var class byte
|
||||
var width int
|
||||
var signed bool
|
||||
var bools []bool
|
||||
var ints []int64
|
||||
var i8s []int8
|
||||
var u8s []uint8
|
||||
var i16s []int16
|
||||
var u16s []uint16
|
||||
var i32s []int32
|
||||
var u32s []uint32
|
||||
var f32s []float32
|
||||
var f64s []float64
|
||||
switch d.Values.Dtype() {
|
||||
case core.Bool:
|
||||
class, width, bools = 8, 1, d.Values.RawBools()
|
||||
case core.Int8:
|
||||
class, width, signed, i8s = 0, 1, true, d.Values.RawInt8s()
|
||||
case core.Uint8:
|
||||
class, width, u8s = 0, 1, d.Values.RawUint8s()
|
||||
case core.Int16:
|
||||
class, width, signed, i16s = 0, 2, true, d.Values.RawInt16s()
|
||||
case core.Uint16:
|
||||
class, width, u16s = 0, 2, d.Values.RawUint16s()
|
||||
case core.Int32:
|
||||
class, width, signed, i32s = 0, 4, true, d.Values.RawInt32s()
|
||||
case core.Uint32:
|
||||
class, width, u32s = 0, 4, d.Values.RawUint32s()
|
||||
case core.Int:
|
||||
class, width, signed, ints = 0, 8, true, d.Values.RawInts()
|
||||
case core.Float32:
|
||||
class, width, f32s = 1, 4, d.Values.RawFloat32s()
|
||||
case core.Float:
|
||||
class, width, f64s = 1, 8, d.Values.RawFloats()
|
||||
default:
|
||||
return nil, base.Errf("%s: dataset %q: dtype %s is not supported; the writer stores bool, int8, uint8, int16, uint16, int32, uint32, int64, float32 and float64", name, d.Path, d.Values.Dtype())
|
||||
}
|
||||
shape := d.Values.Shape()
|
||||
if d.Shape != nil && !slices.Equal(d.Shape, shape) {
|
||||
return nil, base.Errf("%s: dataset %q declares a shape of %v for values shaped %v", name, d.Path, d.Shape, shape)
|
||||
}
|
||||
if len(shape) > hdf5MaxRank {
|
||||
return nil, base.Errf("%s: dataset %q has %d dimensions, the format allows %d", name, d.Path, len(shape), hdf5MaxRank)
|
||||
}
|
||||
// The payload's byte size is bounded here, in the plan, so the write
|
||||
// can reserve it whole: this check is what stands between a wrapped
|
||||
// shape and a reservation past the budget.
|
||||
n, err := hdf5ByteExtent(shape, width, hdf5MaxDatasetBytes)
|
||||
if err != nil {
|
||||
return nil, base.Errf("%s: dataset %q: %w", name, d.Path, err)
|
||||
}
|
||||
attrs, err := hdf5PlanAttrs(name, d.Path, d.Attrs)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &hdf5OutSet{
|
||||
path: d.Path, shape: shape, class: class, width: width, signed: signed,
|
||||
nbytes: int(n), bools: bools, ints: ints, i8s: i8s, u8s: u8s,
|
||||
i16s: i16s, u16s: u16s, i32s: i32s, u32s: u32s,
|
||||
f32s: f32s, f64s: f64s, attrs: attrs,
|
||||
}, nil
|
||||
}
|
||||
|
||||
// hdf5PlanText validates one string dataset and pads its elements to
|
||||
// the longest one, the fixed length the datatype message declares.
|
||||
func hdf5PlanText(name string, tx *HDF5TextDataset) (*hdf5OutSet, error) {
|
||||
if len(tx.Shape) > hdf5MaxRank {
|
||||
return nil, base.Errf("%s: dataset %q has %d dimensions, the format allows %d", name, tx.Path, len(tx.Shape), hdf5MaxRank)
|
||||
}
|
||||
n := 1
|
||||
for _, d := range tx.Shape {
|
||||
if d < 0 {
|
||||
return nil, base.Errf("%s: dataset %q has the negative extent %d", name, tx.Path, d)
|
||||
}
|
||||
n *= d
|
||||
}
|
||||
if len(tx.Text) != n {
|
||||
return nil, base.Errf("%s: dataset %q holds %d strings for a shape of %d elements", name, tx.Path, len(tx.Text), n)
|
||||
}
|
||||
width := 1
|
||||
for _, s := range tx.Text {
|
||||
if strings.IndexByte(s, 0) >= 0 {
|
||||
return nil, base.Errf("%s: dataset %q holds a string with a NUL byte, which a fixed-length element cannot carry", name, tx.Path)
|
||||
}
|
||||
width = max(width, len(s))
|
||||
}
|
||||
// The same byte budget the numeric plan answers to: a shape whose
|
||||
// extents wrap the element count onto len(nil) would otherwise pass
|
||||
// the length check and write a header declaring data it does not
|
||||
// store.
|
||||
if _, err := hdf5ByteExtent(tx.Shape, width, hdf5MaxDatasetBytes); err != nil {
|
||||
return nil, base.Errf("%s: dataset %q: %w", name, tx.Path, err)
|
||||
}
|
||||
// Every element occupies the fixed width; the write lays the
|
||||
// strings into the image itself, where the padding behind each is
|
||||
// the zero the buffer already holds.
|
||||
return &hdf5OutSet{path: tx.Path, shape: tx.Shape, class: 3, width: width, nbytes: n * width, texts: tx.Text}, nil
|
||||
}
|
||||
|
||||
// encode lays the dataset's payload into dst, which holds exactly the
|
||||
// nbytes the plan bounded: every numeric value little-endian at its own
|
||||
// stored width, bool as one zero-or-one byte per element, the strings
|
||||
// each into the fixed-width slot its index names. The zeros dst arrives
|
||||
// with are the padding behind every string, so the encoder writes only
|
||||
// the bytes the values themselves fill. A class or width the plan never
|
||||
// produces is a loud error, never a silent zero payload.
|
||||
func (s *hdf5OutSet) encode(dst []byte) error {
|
||||
switch s.class {
|
||||
case 0:
|
||||
switch {
|
||||
case s.width == 1 && s.signed:
|
||||
for i, v := range s.i8s {
|
||||
dst[i] = byte(v)
|
||||
}
|
||||
case s.width == 1:
|
||||
for i, v := range s.u8s {
|
||||
dst[i] = v
|
||||
}
|
||||
case s.width == 2 && s.signed:
|
||||
for i, v := range s.i16s {
|
||||
binary.LittleEndian.PutUint16(dst[i*2:], uint16(v))
|
||||
}
|
||||
case s.width == 2:
|
||||
for i, v := range s.u16s {
|
||||
binary.LittleEndian.PutUint16(dst[i*2:], v)
|
||||
}
|
||||
case s.width == 4 && s.signed:
|
||||
for i, v := range s.i32s {
|
||||
binary.LittleEndian.PutUint32(dst[i*4:], uint32(v))
|
||||
}
|
||||
case s.width == 4:
|
||||
for i, v := range s.u32s {
|
||||
binary.LittleEndian.PutUint32(dst[i*4:], v)
|
||||
}
|
||||
case s.width == 8:
|
||||
for i, v := range s.ints {
|
||||
binary.LittleEndian.PutUint64(dst[i*8:], uint64(v))
|
||||
}
|
||||
default:
|
||||
return s.payloadRefusal()
|
||||
}
|
||||
case 1:
|
||||
switch s.width {
|
||||
case 4:
|
||||
for i, v := range s.f32s {
|
||||
binary.LittleEndian.PutUint32(dst[i*4:], math.Float32bits(v))
|
||||
}
|
||||
case 8:
|
||||
for i, v := range s.f64s {
|
||||
binary.LittleEndian.PutUint64(dst[i*8:], math.Float64bits(v))
|
||||
}
|
||||
default:
|
||||
return s.payloadRefusal()
|
||||
}
|
||||
case 3:
|
||||
for i, t := range s.texts {
|
||||
copy(dst[i*s.width:], t)
|
||||
}
|
||||
case 8:
|
||||
// The enumeration's member values, one byte per element.
|
||||
for i, v := range s.bools {
|
||||
if v {
|
||||
dst[i] = 1
|
||||
} else {
|
||||
dst[i] = 0
|
||||
}
|
||||
}
|
||||
default:
|
||||
return s.payloadRefusal()
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// payloadRefusal names the datatype an encoder cannot serialise. The
|
||||
// plan produces none of them, so reaching one is a writer defect, and
|
||||
// it fails loudly rather than emitting a payload of silent zeros.
|
||||
func (s *hdf5OutSet) payloadRefusal() error {
|
||||
return base.Errf("dataset %q: the writer cannot serialise datatype class %d of %d bytes per element", s.path, s.class, s.width)
|
||||
}
|
||||
|
||||
// hdf5Attr is one attribute of the plan: its name and the text its
|
||||
// value is written from.
|
||||
type hdf5Attr struct {
|
||||
name string
|
||||
text string
|
||||
}
|
||||
|
||||
// hdf5PlanAttrs validates and sorts the attributes of one object: a
|
||||
// name must be non-empty and free of NUL bytes, the same constraint
|
||||
// the reader's rendering can round-trip under.
|
||||
func hdf5PlanAttrs(name, path string, attrs map[string]string) ([]hdf5Attr, error) {
|
||||
if len(attrs) > hdf5MaxAttrs {
|
||||
return nil, base.Errf("%s: %q carries %d attributes, past the %d the writer stores in one object header", name, path, len(attrs), hdf5MaxAttrs)
|
||||
}
|
||||
keys := slices.Sorted(maps.Keys(attrs))
|
||||
out := make([]hdf5Attr, 0, len(keys))
|
||||
for _, k := range keys {
|
||||
if k == "" {
|
||||
return nil, base.Errf("%s: %q carries an attribute with an empty name", name, path)
|
||||
}
|
||||
if len(k)+1 > 0xffff {
|
||||
return nil, base.Errf("%s: %q carries the attribute %q whose name is past the %d bytes the attribute message counts", name, path, k, 0xffff)
|
||||
}
|
||||
if strings.IndexByte(k, 0) >= 0 || strings.IndexByte(attrs[k], 0) >= 0 {
|
||||
return nil, base.Errf("%s: %q carries the attribute %q with a NUL byte in its name or value", name, path, k)
|
||||
}
|
||||
out = append(out, hdf5Attr{name: k, text: attrs[k]})
|
||||
}
|
||||
return out, nil
|
||||
}
|
||||
|
||||
// hdf5Writer builds the file image: every address is an offset into
|
||||
// buf, so structures written later can be referenced by structures
|
||||
// written earlier through the patch at the end. The chunk staging
|
||||
// fields are reused across every chunk of one write: the gather and
|
||||
// shuffle buffers and the index vectors grow to the largest chunk the
|
||||
// write lays out, and the deflater carries one compressor and one
|
||||
// output buffer for the whole file.
|
||||
type hdf5Writer struct {
|
||||
buf []byte
|
||||
latest bool
|
||||
opts HDF5WriteOptions
|
||||
sbAddr uint64
|
||||
|
||||
gatherScratch []byte
|
||||
shuffleScratch []byte
|
||||
chunkIdx []int
|
||||
comp *zlib.Writer
|
||||
compBuf bytes.Buffer
|
||||
compLevel int
|
||||
}
|
||||
|
||||
// write lays the file out the way the reference library builds it: a
|
||||
// placeholder superblock first (its fields name the end of the file
|
||||
// and the root group, which are known only once everything is
|
||||
// written), then the root group, whose header is allocated before its
|
||||
// subtree and filled once the subtree has addresses, and the
|
||||
// superblock itself last.
|
||||
func (w *hdf5Writer) write(root *hdf5OutNode) error {
|
||||
// The image is built into one buffer, so it is allocated once from
|
||||
// the plan's own size: growing it as the structures are laid out
|
||||
// would copy the whole file at every step.
|
||||
w.buf = make([]byte, 0, hdf5ImageEstimate(root, w.opts))
|
||||
if w.latest {
|
||||
w.sbAddr = w.alloc(hdf5Superblock3Size)
|
||||
} else {
|
||||
w.sbAddr = w.alloc(hdf5Superblock0Size)
|
||||
}
|
||||
if err := w.writeGroup(root); err != nil {
|
||||
return err
|
||||
}
|
||||
w.finishSuperblock(root)
|
||||
return nil
|
||||
}
|
||||
|
||||
// writeGroup dispatches to the group writer of the file's layout.
|
||||
func (w *hdf5Writer) writeGroup(g *hdf5OutNode) error {
|
||||
if w.latest {
|
||||
return w.writeNewGroup(g)
|
||||
}
|
||||
return w.writeClassicGroup(g)
|
||||
}
|
||||
|
||||
// Sizes of the fixed parts the writer places first. Every address and
|
||||
// length is eight bytes, as in the fixtures.
|
||||
const (
|
||||
hdf5Superblock0Size = 24 + 4*8 + 8 + 8 + 4 + 4 + 16 // 96
|
||||
hdf5Superblock3Size = 12 + 4*8 + 4 // 48
|
||||
)
|
||||
|
||||
// finishSuperblock fills the placeholder: the classic superblock
|
||||
// names the root group through a symbol table entry whose cache holds
|
||||
// the group's B-tree and local heap, the latest one names its object
|
||||
// header directly and checksums the whole block with lookup3.
|
||||
func (w *hdf5Writer) finishSuperblock(root *hdf5OutNode) {
|
||||
copy(w.buf[w.sbAddr:], hdf5Magic)
|
||||
if w.latest {
|
||||
w.buf[w.sbAddr+8] = 3
|
||||
w.buf[w.sbAddr+9] = 8
|
||||
w.buf[w.sbAddr+10] = 8
|
||||
w.buf[w.sbAddr+11] = 0 // file consistency flags
|
||||
w.set64(w.sbAddr+12, 0)
|
||||
w.set64(w.sbAddr+20, math.MaxUint64)
|
||||
w.set64(w.sbAddr+28, uint64(len(w.buf)))
|
||||
w.set64(w.sbAddr+36, root.addr)
|
||||
w.set32(w.sbAddr+44, hdf5Lookup3(w.buf[w.sbAddr:w.sbAddr+44]))
|
||||
return
|
||||
}
|
||||
w.buf[w.sbAddr+8] = 0 // superblock version
|
||||
w.buf[w.sbAddr+9] = 0 // free space storage version
|
||||
w.buf[w.sbAddr+10] = 0 // root group symbol table entry version
|
||||
w.buf[w.sbAddr+11] = 0 // reserved
|
||||
w.buf[w.sbAddr+12] = 0 // shared header message format version
|
||||
w.buf[w.sbAddr+13] = 8 // size of offsets
|
||||
w.buf[w.sbAddr+14] = 8 // size of lengths
|
||||
w.buf[w.sbAddr+15] = 0 // reserved
|
||||
binary.LittleEndian.PutUint16(w.buf[w.sbAddr+16:], hdf5GroupLeafK)
|
||||
binary.LittleEndian.PutUint16(w.buf[w.sbAddr+18:], hdf5GroupInnerK)
|
||||
// The file consistency flags at +20 stay zero.
|
||||
w.set64(w.sbAddr+24, 0) // base address
|
||||
w.set64(w.sbAddr+32, math.MaxUint64) // free space information
|
||||
w.set64(w.sbAddr+40, uint64(len(w.buf))) // end of file
|
||||
w.set64(w.sbAddr+48, math.MaxUint64) // driver information
|
||||
w.set64(w.sbAddr+56, 0) // root entry: link name offset
|
||||
w.set64(w.sbAddr+64, root.addr) // root entry: object header address
|
||||
w.set32(w.sbAddr+72, 1) // root entry: symbol table cache
|
||||
w.set64(w.sbAddr+80, root.btree) // cache: B-tree address
|
||||
w.set64(w.sbAddr+88, root.heap) // cache: local heap address
|
||||
}
|
||||
|
||||
func (w *hdf5Writer) alloc(n int) uint64 {
|
||||
addr := uint64(len(w.buf))
|
||||
w.buf = append(w.buf, make([]byte, n)...)
|
||||
return addr
|
||||
}
|
||||
|
||||
// reserve extends the image by n bytes without writing them and
|
||||
// returns the address they start at; the caller fills the whole span
|
||||
// in the same breath, so the payload passes through the writer once.
|
||||
// The capacity beyond len(buf) always holds the zeros the buffer's
|
||||
// allocations left there, which every fixed structure and string slot
|
||||
// is padded from, and the encode that follows a reservation writes
|
||||
// every byte the payload itself does not.
|
||||
func (w *hdf5Writer) reserve(n int) uint64 {
|
||||
addr := uint64(len(w.buf))
|
||||
if cap(w.buf)-len(w.buf) < n {
|
||||
w.buf = append(w.buf, make([]byte, n)...)
|
||||
return addr
|
||||
}
|
||||
w.buf = w.buf[:len(w.buf)+n]
|
||||
return addr
|
||||
}
|
||||
|
||||
func (w *hdf5Writer) bytes(b []byte) uint64 {
|
||||
addr := uint64(len(w.buf))
|
||||
w.buf = append(w.buf, b...)
|
||||
return addr
|
||||
}
|
||||
|
||||
// pad8 aligns the image to the eight-byte boundary the format inserts
|
||||
// between the structures of the classic layout.
|
||||
func (w *hdf5Writer) pad8() {
|
||||
if r := len(w.buf) % 8; r != 0 {
|
||||
w.buf = append(w.buf, make([]byte, 8-r)...)
|
||||
}
|
||||
}
|
||||
|
||||
func (w *hdf5Writer) set32(at uint64, v uint32) {
|
||||
binary.LittleEndian.PutUint32(w.buf[at:], v)
|
||||
}
|
||||
|
||||
func (w *hdf5Writer) set64(at uint64, v uint64) {
|
||||
binary.LittleEndian.PutUint64(w.buf[at:], v)
|
||||
}
|
||||
@@ -0,0 +1,23 @@
|
||||
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
|
||||
// SPDX-License-Identifier: MIT
|
||||
|
||||
package io
|
||||
|
||||
import "testing"
|
||||
|
||||
// TestHDF5SetEstimateCapsTinyChunkTargets pins the node estimate's
|
||||
// cap: the estimate is a starting buffer the write grows past by
|
||||
// append, so a pathological chunk target of a couple of bytes must
|
||||
// charge the dataset a bounded guess rather than a B-tree node, and
|
||||
// its kilobyte of image, to every single data byte.
|
||||
func TestHDF5SetEstimateCapsTinyChunkTargets(t *testing.T) {
|
||||
s := &hdf5OutSet{path: "d", shape: []int{1 << 20}, class: 1, width: 8, nbytes: 8 << 20}
|
||||
honest := hdf5SetEstimate(s, HDF5WriteOptions{ChunkBytes: 4096, Shuffle: true})
|
||||
if honest <= int64(s.nbytes) || honest > 100<<20 {
|
||||
t.Fatalf("the honest estimate is %d bytes for a %d-byte payload, want a bounded cover", honest, s.nbytes)
|
||||
}
|
||||
tiny := hdf5SetEstimate(s, HDF5WriteOptions{ChunkBytes: 2, Shuffle: true})
|
||||
if tiny > 50<<20 {
|
||||
t.Fatalf("a two-byte chunk target estimated %d bytes for a %d-byte payload", tiny, s.nbytes)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,717 @@
|
||||
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
|
||||
// SPDX-License-Identifier: MIT
|
||||
|
||||
package io
|
||||
|
||||
// Group and chunked-storage writing for the HDF5 writer. Both group
|
||||
// layouts the reader accepts are written: the classic symbol table (a
|
||||
// local heap of names, symbol table nodes and a version 1 group
|
||||
// B-tree) and the latest-style compact group (link messages in the
|
||||
// object header). Chunked storage hangs its chunks off a version 1
|
||||
// chunk B-tree, whose conventions the reference files pin: every node
|
||||
// is allocated at the size the format's fanout implies, an interior
|
||||
// node's key is the first key of the child's subtree, and a node's
|
||||
// closing key is a sentinel ordered past its last chunk.
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"compress/zlib"
|
||||
"encoding/binary"
|
||||
"math"
|
||||
"slices"
|
||||
"strings"
|
||||
|
||||
"sourcedock.dev/petrbalvin/tensor/internal/base"
|
||||
)
|
||||
|
||||
// hdf5GroupKid is one child pointer of a group B-tree node. The key of
|
||||
// kid i is the heap offset of the last name child i-1 holds, so a
|
||||
// node's entries are headed by the zero key (the heap's null name,
|
||||
// below every real name) and each entry's key closes the range of the
|
||||
// child before it, which is the convention the fixtures store.
|
||||
type hdf5GroupKid struct {
|
||||
key uint64
|
||||
child uint64
|
||||
}
|
||||
|
||||
// hdf5GroupChild is one child entry of a group: its link name, its
|
||||
// object header address and, until the children are written, the
|
||||
// subgroup or dataset whose address it carries.
|
||||
type hdf5GroupChild struct {
|
||||
name string
|
||||
addr uint64
|
||||
group *hdf5OutNode
|
||||
set *hdf5OutSet
|
||||
}
|
||||
|
||||
// hdf5Kids merges a group's subgroups and datasets into one list in
|
||||
// name order, the order the classic symbol table requires: the local
|
||||
// heap assigns its offsets in list order and the reference library
|
||||
// binary-searches both the B-tree keys and the symbol node records on
|
||||
// those offsets, so a groups-first order would hide every child whose
|
||||
// name interleaves with a dataset's. The latest layout's link list
|
||||
// keeps the same order.
|
||||
func hdf5Kids(g *hdf5OutNode) []hdf5GroupChild {
|
||||
kids := make([]hdf5GroupChild, 0, len(g.groups)+len(g.sets))
|
||||
for _, sub := range g.groups {
|
||||
kids = append(kids, hdf5GroupChild{name: sub.name, group: sub})
|
||||
}
|
||||
for _, s := range g.sets {
|
||||
kids = append(kids, hdf5GroupChild{name: baseName(s.path), set: s})
|
||||
}
|
||||
slices.SortFunc(kids, func(a, b hdf5GroupChild) int { return strings.Compare(a.name, b.name) })
|
||||
return kids
|
||||
}
|
||||
|
||||
// kidAddr is a child's written object header address.
|
||||
func kidAddr(k hdf5GroupChild) uint64 {
|
||||
if k.group != nil {
|
||||
return k.group.addr
|
||||
}
|
||||
return k.set.addr
|
||||
}
|
||||
|
||||
// writeClassicGroup writes a group in the classic layout, in the
|
||||
// construction order the reference library uses: the object header is
|
||||
// allocated first (over a placeholder, since the symbol table message
|
||||
// can only name the B-tree and heap once they exist), then the local
|
||||
// heap, then the children, and the symbol table nodes and B-tree last.
|
||||
// A child group's symbol entry carries the B-tree and heap addresses
|
||||
// in its cache, as the reference writes; a dataset's carries none.
|
||||
func (w *hdf5Writer) writeClassicGroup(g *hdf5OutNode) error {
|
||||
kids := hdf5Kids(g)
|
||||
// The header placeholder: sized and shaped exactly like the final
|
||||
// image, with the B-tree and heap addresses still zero.
|
||||
msgs, err := w.classicGroupMsgs(g, 0, 0)
|
||||
if err != nil {
|
||||
return base.Errf("SaveHDF5: group %q: %w", g.path, err)
|
||||
}
|
||||
image, err := hdf5HeaderV1(msgs)
|
||||
if err != nil {
|
||||
return base.Errf("SaveHDF5: group %q: %w", g.path, err)
|
||||
}
|
||||
hdrAddr := w.bytes(image)
|
||||
// The heap assigns every name its offset; the children are in
|
||||
// sorted order, so offset order and name order agree and the
|
||||
// B-tree's byte-offset keys sort as the names do. The heap carries
|
||||
// slack behind the names, written as the free block the reference
|
||||
// library leaves there: a next offset of one is the sentinel that
|
||||
// ends the free list.
|
||||
heapData := make([]byte, 8) // offset 0 is the null name no entry points at
|
||||
offsets := make([]uint64, len(kids))
|
||||
for i, k := range kids {
|
||||
offsets[i] = uint64(len(heapData))
|
||||
heapData = append(heapData, k.name...)
|
||||
heapData = hdf5AppendAlign(append(heapData, 0))
|
||||
}
|
||||
size := max(len(heapData)+16, 88)
|
||||
free := len(heapData)
|
||||
heapData = append(heapData, make([]byte, 16)...) // the free block: next, size
|
||||
binary.LittleEndian.PutUint64(heapData[free:], 1)
|
||||
binary.LittleEndian.PutUint64(heapData[free+8:], uint64(size-free))
|
||||
heapData = append(heapData, make([]byte, size-len(heapData))...)
|
||||
dataAddr := w.bytes(heapData)
|
||||
w.pad8()
|
||||
heap := w.writeLocalHeap(dataAddr, size, free)
|
||||
// The children: datasets, then the subgroups, each of which lays
|
||||
// out its own subtree in the children's name order.
|
||||
for _, s := range g.sets {
|
||||
addr, err := w.writeDataset(s)
|
||||
if err != nil {
|
||||
// writeDataset names the dataset itself; a second wrap
|
||||
// here would prefix it twice.
|
||||
return err
|
||||
}
|
||||
s.addr = addr
|
||||
}
|
||||
for _, sub := range g.groups {
|
||||
if err := w.writeGroup(sub); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
for i := range kids {
|
||||
kids[i].addr = kidAddr(kids[i])
|
||||
}
|
||||
// Symbol table nodes, eight entries each: the node occupies the
|
||||
// size the format's leaf K of four implies, however few entries are
|
||||
// used, and an empty group still holds one.
|
||||
leaves := make([]hdf5GroupKid, 0, max((len(kids)+2*hdf5GroupLeafK-1)/(2*hdf5GroupLeafK), 1))
|
||||
for start := 0; start < len(kids); start += 2 * hdf5GroupLeafK {
|
||||
leaves = append(leaves, w.writeSymbolNode(kids[start:min(start+2*hdf5GroupLeafK, len(kids))], offsets[start:]))
|
||||
}
|
||||
if len(leaves) == 0 {
|
||||
leaves = append(leaves, w.writeSymbolNode(nil, nil))
|
||||
}
|
||||
btree := w.writeGroupTree(leaves)
|
||||
// The header over the placeholder, now that the B-tree and heap
|
||||
// have their final addresses.
|
||||
msgs, err = w.classicGroupMsgs(g, btree, heap)
|
||||
if err != nil {
|
||||
return base.Errf("SaveHDF5: group %q: %w", g.path, err)
|
||||
}
|
||||
image, err = hdf5HeaderV1(msgs)
|
||||
if err != nil {
|
||||
return base.Errf("SaveHDF5: group %q: %w", g.path, err)
|
||||
}
|
||||
w.headerAt(hdrAddr, image)
|
||||
g.addr, g.btree, g.heap = hdrAddr, btree, heap
|
||||
return nil
|
||||
}
|
||||
|
||||
// classicGroupMsgs builds the message list of a classic group's object
|
||||
// header: the symbol table message naming the B-tree and heap, then
|
||||
// the attributes in name order.
|
||||
func (w *hdf5Writer) classicGroupMsgs(g *hdf5OutNode, btree, heap uint64) ([]hdf5OutMsg, error) {
|
||||
msgs := []hdf5OutMsg{{typ: hdf5MsgSymbolTable, data: w.appendAddrs(nil, btree, heap)}}
|
||||
for _, a := range g.attrs {
|
||||
m, err := hdf5AttrMessage(a)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
msgs = append(msgs, m)
|
||||
}
|
||||
return msgs, nil
|
||||
}
|
||||
|
||||
// writeSymbolNode writes one symbol table node of the given entries,
|
||||
// allocated at eight slots. A child group's cache holds its B-tree and
|
||||
// heap addresses; a dataset's cache stays empty.
|
||||
func (w *hdf5Writer) writeSymbolNode(batch []hdf5GroupChild, offsets []uint64) hdf5GroupKid {
|
||||
node := make([]byte, 8+2*hdf5GroupLeafK*(8+8+24))
|
||||
copy(node, hdf5SymbolNode)
|
||||
node[4] = 1
|
||||
binary.LittleEndian.PutUint16(node[6:], uint16(len(batch)))
|
||||
p := 8
|
||||
for i, k := range batch {
|
||||
binary.LittleEndian.PutUint64(node[p:], offsets[i])
|
||||
binary.LittleEndian.PutUint64(node[p+8:], k.addr)
|
||||
if k.group != nil {
|
||||
binary.LittleEndian.PutUint32(node[p+16:], 1)
|
||||
binary.LittleEndian.PutUint64(node[p+24:], k.group.btree)
|
||||
binary.LittleEndian.PutUint64(node[p+32:], k.group.heap)
|
||||
}
|
||||
p += 8 + 8 + 24
|
||||
}
|
||||
leaf := hdf5GroupKid{child: w.bytes(node)}
|
||||
if len(batch) > 0 {
|
||||
leaf.key = offsets[len(batch)-1]
|
||||
}
|
||||
return leaf
|
||||
}
|
||||
|
||||
// baseName returns the last segment of an absolute path.
|
||||
func baseName(path string) string {
|
||||
for i := len(path) - 1; i >= 0; i-- {
|
||||
if path[i] == '/' {
|
||||
return path[i+1:]
|
||||
}
|
||||
}
|
||||
return path
|
||||
}
|
||||
|
||||
// writeNewGroup writes a group in the latest layout, the construction
|
||||
// order the reference library uses: the object header is allocated
|
||||
// first over a placeholder, then the children, and the header is
|
||||
// written over the placeholder once every link's address is known.
|
||||
// The link info and group info messages the reference writes head the
|
||||
// message list; a group without links carries no link info message:
|
||||
// beside no links the reader of this package takes that message for
|
||||
// dense storage, which it refuses.
|
||||
func (w *hdf5Writer) writeNewGroup(g *hdf5OutNode) error {
|
||||
msgs, err := w.newGroupMsgs(g)
|
||||
if err != nil {
|
||||
return base.Errf("SaveHDF5: group %q: %w", g.path, err)
|
||||
}
|
||||
image, err := hdf5HeaderV2(msgs)
|
||||
if err != nil {
|
||||
return base.Errf("SaveHDF5: group %q: %w", g.path, err)
|
||||
}
|
||||
hdrAddr := w.bytes(image)
|
||||
for _, s := range g.sets {
|
||||
addr, err := w.writeDataset(s)
|
||||
if err != nil {
|
||||
// writeDataset names the dataset itself; a second wrap
|
||||
// here would prefix it twice.
|
||||
return err
|
||||
}
|
||||
s.addr = addr
|
||||
}
|
||||
for _, sub := range g.groups {
|
||||
if err := w.writeGroup(sub); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
msgs, err = w.newGroupMsgs(g)
|
||||
if err != nil {
|
||||
return base.Errf("SaveHDF5: group %q: %w", g.path, err)
|
||||
}
|
||||
image, err = hdf5HeaderV2(msgs)
|
||||
if err != nil {
|
||||
return base.Errf("SaveHDF5: group %q: %w", g.path, err)
|
||||
}
|
||||
w.headerAt(hdrAddr, image)
|
||||
g.addr = hdrAddr
|
||||
return nil
|
||||
}
|
||||
|
||||
// newGroupMsgs builds the message list of a latest-style group's
|
||||
// object header. Addresses unknown at placeholder time are zero; the
|
||||
// message sizes are the same either way.
|
||||
func (w *hdf5Writer) newGroupMsgs(g *hdf5OutNode) ([]hdf5OutMsg, error) {
|
||||
var msgs []hdf5OutMsg
|
||||
if len(g.groups)+len(g.sets) > 0 {
|
||||
undef := bytes.Repeat([]byte{0xff}, 8)
|
||||
info := append([]byte{0, 0}, undef...)
|
||||
info = append(info, undef...)
|
||||
msgs = append(msgs,
|
||||
hdf5OutMsg{typ: hdf5MsgLinkInfo, data: info},
|
||||
hdf5OutMsg{typ: hdf5MsgGroupInfo, data: []byte{0, 0}},
|
||||
)
|
||||
}
|
||||
addLink := func(name string, addr uint64) {
|
||||
body := []byte{1, 0} // version 1, no creation order, one-byte length
|
||||
if len(name) >= 256 {
|
||||
body[1] = 0x01 // the name length widens to two bytes
|
||||
body = binary.LittleEndian.AppendUint16(body, uint16(len(name)))
|
||||
} else {
|
||||
body = append(body, byte(len(name)))
|
||||
}
|
||||
body = append(body, name...)
|
||||
body = binary.LittleEndian.AppendUint64(body, addr)
|
||||
msgs = append(msgs, hdf5OutMsg{typ: hdf5MsgLink, data: body})
|
||||
}
|
||||
for _, k := range hdf5Kids(g) {
|
||||
addLink(k.name, kidAddr(k))
|
||||
}
|
||||
for _, a := range g.attrs {
|
||||
m, err := hdf5AttrMessage(a)
|
||||
if err != nil {
|
||||
// Unwrapped: the caller carries the group context and
|
||||
// prefixes the entry point, so wrapping here would double
|
||||
// both.
|
||||
return nil, err
|
||||
}
|
||||
msgs = append(msgs, m)
|
||||
}
|
||||
return msgs, nil
|
||||
}
|
||||
|
||||
// writeLocalHeap writes a local heap header in front of the data
|
||||
// segment written earlier: the data segment's size, the free list head
|
||||
// and the segment address.
|
||||
func (w *hdf5Writer) writeLocalHeap(dataAddr uint64, size, free int) uint64 {
|
||||
head := append([]byte{}, hdf5LocalHeap...)
|
||||
head = append(head, 0, 0, 0, 0) // version and reserved
|
||||
head = binary.LittleEndian.AppendUint64(head, uint64(size))
|
||||
head = binary.LittleEndian.AppendUint64(head, uint64(free))
|
||||
head = binary.LittleEndian.AppendUint64(head, dataAddr)
|
||||
return w.bytes(head)
|
||||
}
|
||||
|
||||
// writeGroupTree packs the symbol table nodes into version 1 group
|
||||
// B-tree nodes of thirty-two children and returns the root's address.
|
||||
func (w *hdf5Writer) writeGroupTree(leaves []hdf5GroupKid) uint64 {
|
||||
kids := leaves
|
||||
level := byte(0)
|
||||
for len(kids) > 2*hdf5GroupInnerK {
|
||||
var next []hdf5GroupKid
|
||||
for start := 0; start < len(kids); start += 2 * hdf5GroupInnerK {
|
||||
batch := kids[start:min(start+2*hdf5GroupInnerK, len(kids))]
|
||||
addr := w.writeGroupNode(level, batch)
|
||||
next = append(next, hdf5GroupKid{key: batch[len(batch)-1].key, child: addr})
|
||||
}
|
||||
kids = next
|
||||
level++
|
||||
}
|
||||
return w.writeGroupNode(level, kids)
|
||||
}
|
||||
|
||||
// writeGroupNode writes one group B-tree node, allocated at the size
|
||||
// the format's internal K of sixteen implies. The sibling addresses
|
||||
// are the undefined value: a zero would be a defined address, and a
|
||||
// reader would follow it as a sibling node.
|
||||
func (w *hdf5Writer) writeGroupNode(level byte, kids []hdf5GroupKid) uint64 {
|
||||
node := make([]byte, 8+2*8+2*hdf5GroupInnerK*(8+8)+8)
|
||||
copy(node, hdf5Tree)
|
||||
node[4] = 0 // a group B-tree
|
||||
node[5] = level
|
||||
binary.LittleEndian.PutUint16(node[6:], uint16(len(kids)))
|
||||
binary.LittleEndian.PutUint64(node[8:], math.MaxUint64) // left sibling
|
||||
binary.LittleEndian.PutUint64(node[16:], math.MaxUint64) // right sibling
|
||||
p := 24
|
||||
binary.LittleEndian.PutUint64(node[p:], 0) // key[0]: the null name
|
||||
p += 8
|
||||
for _, k := range kids {
|
||||
binary.LittleEndian.PutUint64(node[p:], k.child)
|
||||
p += 8
|
||||
binary.LittleEndian.PutUint64(node[p:], k.key) // the key closing this child
|
||||
p += 8
|
||||
}
|
||||
return w.bytes(node)
|
||||
}
|
||||
|
||||
// appendAddrs appends the given addresses as the file's eight-byte
|
||||
// address fields.
|
||||
func (w *hdf5Writer) appendAddrs(b []byte, addrs ...uint64) []byte {
|
||||
for _, a := range addrs {
|
||||
b = binary.LittleEndian.AppendUint64(b, a)
|
||||
}
|
||||
return b
|
||||
}
|
||||
|
||||
// hdf5ChunkKey is a version 1 chunk B-tree key: the chunk's stored
|
||||
// size, its coordinates, one per dataset dimension plus the element
|
||||
// size dimension the format keys last, and the address the chunk's
|
||||
// bytes occupy.
|
||||
type hdf5ChunkKey struct {
|
||||
size uint32
|
||||
coords []uint64
|
||||
addr uint64
|
||||
}
|
||||
|
||||
// hdf5ChunkKid is one child of a chunk B-tree node with the first and
|
||||
// last keys of its subtree, which become the node's own keys.
|
||||
type hdf5ChunkKid struct {
|
||||
first, last hdf5ChunkKey
|
||||
child uint64
|
||||
}
|
||||
|
||||
// hdf5ChunkNodeSize is the size a chunk B-tree node occupies: the
|
||||
// header, twice the storage K of thirty-two keys with their child
|
||||
// pointers, and the closing key.
|
||||
func hdf5ChunkNodeSize(rank int) int {
|
||||
keySize := 8 + 8*(rank+1)
|
||||
return 8 + 2*8 + 2*hdf5IStoreK*(keySize+8) + keySize
|
||||
}
|
||||
|
||||
// hdf5ChunkSentinel is the key that closes a node: the last chunk's
|
||||
// coordinates with the element size dimension set to the element size,
|
||||
// which orders it past every chunk of the node, as the reference
|
||||
// files store it. Reference tooling reads either shape of the
|
||||
// sentinel: recent releases write the chunk shape itself with the
|
||||
// element size appended, which orders identically; this writer keeps
|
||||
// the fixture convention.
|
||||
func hdf5ChunkSentinel(k hdf5ChunkKey, width int) hdf5ChunkKey {
|
||||
coords := append([]uint64{}, k.coords...)
|
||||
coords[len(coords)-1] = uint64(width)
|
||||
return hdf5ChunkKey{size: 0, coords: coords}
|
||||
}
|
||||
|
||||
// writeChunks chunks one dataset, applies the pipeline to every chunk
|
||||
// and lays the chunks out through a chunk B-tree, returning the
|
||||
// B-tree's address and the chunk shape. A chunk is stored complete:
|
||||
// the edge chunks of a dataset whose extent is not a whole number of
|
||||
// chunks are padded with the zero fill value, which is what chunked
|
||||
// storage holds and what the reader requires. The staging buffers, the
|
||||
// index vectors and the deflater are writer state reused across every
|
||||
// chunk and every dataset of the write; each chunk fully overwrites
|
||||
// what it stages, so no bytes travel between chunks.
|
||||
func (w *hdf5Writer) writeChunks(s *hdf5OutSet) (uint64, []int, error) {
|
||||
chunk := hdf5ChunkShape(s.shape, s.width, w.chunkTarget())
|
||||
grid := make([]int, len(s.shape))
|
||||
total := 1
|
||||
for i := range grid {
|
||||
grid[i] = (s.shape[i] + chunk[i] - 1) / chunk[i]
|
||||
total *= grid[i]
|
||||
}
|
||||
if total > hdf5MaxChunks {
|
||||
return 0, nil, base.Errf("chunking to a %d byte target needs %d chunks, past the %d the writer lays out; raise the chunk target", w.chunkTarget(), total, hdf5MaxChunks)
|
||||
}
|
||||
if _, err := hdf5ByteExtent(chunk, s.width, hdf5MaxDatasetBytes); err != nil {
|
||||
return 0, nil, base.Errf("the chunk shape %v: %w", chunk, err)
|
||||
}
|
||||
level := w.opts.Gzip
|
||||
if level == -1 {
|
||||
level = 6
|
||||
}
|
||||
keys := make([]hdf5ChunkKey, 0, total)
|
||||
coords := make([]int, len(grid))
|
||||
// Every key aliases one range of a single coordinate slab, which
|
||||
// the chunk B-tree reads before the write ends.
|
||||
coordSlab := make([]uint64, total*(len(chunk)+1))
|
||||
rank := len(s.shape)
|
||||
chunkElems := 1
|
||||
for _, c := range chunk {
|
||||
chunkElems *= c
|
||||
}
|
||||
chunkBytes := chunkElems * s.width
|
||||
if cap(w.gatherScratch) < chunkBytes {
|
||||
w.gatherScratch = make([]byte, chunkBytes)
|
||||
}
|
||||
if w.opts.Shuffle && s.width > 1 && cap(w.shuffleScratch) < chunkBytes {
|
||||
w.shuffleScratch = make([]byte, chunkBytes)
|
||||
}
|
||||
if cap(w.chunkIdx) < 4*rank {
|
||||
w.chunkIdx = make([]int, 4*rank)
|
||||
}
|
||||
origin := w.chunkIdx[:rank]
|
||||
count := w.chunkIdx[rank : 2*rank]
|
||||
srcStride := w.chunkIdx[2*rank : 3*rank]
|
||||
dstStride := w.chunkIdx[3*rank : 4*rank]
|
||||
cs, ds := 1, 1
|
||||
for i := rank - 1; i >= 0; i-- {
|
||||
srcStride[i], dstStride[i] = cs, ds
|
||||
cs *= s.shape[i]
|
||||
ds *= chunk[i]
|
||||
}
|
||||
for ci := range total {
|
||||
for d := range rank {
|
||||
origin[d] = coords[d] * chunk[d]
|
||||
}
|
||||
data := w.gatherScratch[:chunkBytes]
|
||||
if err := w.fillChunk(data, s, chunk, origin, count, srcStride, dstStride); err != nil {
|
||||
// The dataset caller wraps with the dataset's own context;
|
||||
// the bare encode error keeps it single.
|
||||
return 0, nil, err
|
||||
}
|
||||
if w.opts.Shuffle && s.width > 1 {
|
||||
hdf5ShuffleInto(w.shuffleScratch[:chunkBytes], data, s.width)
|
||||
data = w.shuffleScratch[:chunkBytes]
|
||||
}
|
||||
if w.opts.Gzip != 0 {
|
||||
compressed, err := w.deflateChunk(data, level)
|
||||
if err != nil {
|
||||
// The dataset caller wraps with the dataset's own
|
||||
// context; the bare filter error keeps it single.
|
||||
return 0, nil, err
|
||||
}
|
||||
data = compressed
|
||||
}
|
||||
addr := w.bytes(data)
|
||||
w.pad8()
|
||||
// The key's trailing coordinate is the element size dimension
|
||||
// the format keys last, and the reference library carries 0 in
|
||||
// it for every real chunk: only the closing sentinel holds the
|
||||
// element size, which is what orders it past the node's chunks.
|
||||
// A real key carrying the size would tie the sentinel and hide
|
||||
// every chunk from a binary search.
|
||||
key := hdf5ChunkKey{size: uint32(len(data)), coords: coordSlab[ci*(len(chunk)+1) : (ci+1)*(len(chunk)+1)], addr: addr}
|
||||
for i := range chunk {
|
||||
key.coords[i] = uint64(coords[i] * chunk[i])
|
||||
}
|
||||
keys = append(keys, key)
|
||||
// The odometer walks the chunk grid in row-major order, which
|
||||
// is the order the B-tree's keys must be sorted in.
|
||||
for d := len(coords) - 1; d >= 0; d-- {
|
||||
coords[d]++
|
||||
if coords[d] < grid[d] {
|
||||
break
|
||||
}
|
||||
coords[d] = 0
|
||||
}
|
||||
}
|
||||
return w.writeChunkTree(keys, len(s.shape), s.width), chunk, nil
|
||||
}
|
||||
|
||||
// fillChunk stages one chunk of the dataset into dst, which holds the
|
||||
// chunk's full element extent: the cells inside the shape take the
|
||||
// values encoded little-endian at their stored width, exactly the
|
||||
// bytes encode produces from the same values, and the cells past any
|
||||
// edge stay the zeros clear left. Every byte of dst is written on
|
||||
// every call, so the reused gather buffer carries nothing between
|
||||
// chunks. origin is the chunk's first cell in dataset coordinates,
|
||||
// count walks the overlap of the chunk and the shape, and the strides
|
||||
// are row-major element strides of the dataset and of the chunk. A
|
||||
// class or width the plan never produces is the same loud refusal
|
||||
// encode gives.
|
||||
func (w *hdf5Writer) fillChunk(dst []byte, s *hdf5OutSet, chunk, origin, count, srcStride, dstStride []int) error {
|
||||
clear(dst)
|
||||
shape := s.shape
|
||||
rank := len(shape)
|
||||
width := s.width
|
||||
for d := range rank {
|
||||
count[d] = origin[d]
|
||||
}
|
||||
for {
|
||||
src, loc := 0, 0
|
||||
for d := range rank {
|
||||
src += count[d] * srcStride[d]
|
||||
loc += (count[d] - origin[d]) * dstStride[d]
|
||||
}
|
||||
p := loc * width
|
||||
switch {
|
||||
case s.class == 0 && width == 1 && s.signed:
|
||||
dst[p] = byte(s.i8s[src])
|
||||
case s.class == 0 && width == 1:
|
||||
dst[p] = s.u8s[src]
|
||||
case s.class == 0 && width == 2 && s.signed:
|
||||
binary.LittleEndian.PutUint16(dst[p:], uint16(s.i16s[src]))
|
||||
case s.class == 0 && width == 2:
|
||||
binary.LittleEndian.PutUint16(dst[p:], s.u16s[src])
|
||||
case s.class == 0 && width == 4 && s.signed:
|
||||
binary.LittleEndian.PutUint32(dst[p:], uint32(s.i32s[src]))
|
||||
case s.class == 0 && width == 4:
|
||||
binary.LittleEndian.PutUint32(dst[p:], s.u32s[src])
|
||||
case s.class == 0 && width == 8:
|
||||
binary.LittleEndian.PutUint64(dst[p:], uint64(s.ints[src]))
|
||||
case s.class == 8:
|
||||
if s.bools[src] {
|
||||
dst[p] = 1
|
||||
} else {
|
||||
dst[p] = 0
|
||||
}
|
||||
case s.class == 1 && width == 4:
|
||||
binary.LittleEndian.PutUint32(dst[p:], math.Float32bits(s.f32s[src]))
|
||||
case s.class == 1 && width == 8:
|
||||
binary.LittleEndian.PutUint64(dst[p:], math.Float64bits(s.f64s[src]))
|
||||
default:
|
||||
return s.payloadRefusal()
|
||||
}
|
||||
d := rank - 1
|
||||
for d >= 0 {
|
||||
count[d]++
|
||||
if count[d] < min(origin[d]+chunk[d], shape[d]) {
|
||||
break
|
||||
}
|
||||
count[d] = origin[d]
|
||||
d--
|
||||
}
|
||||
if d < 0 {
|
||||
return nil
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// chunkTarget is the configured chunk byte target, or the default.
|
||||
func (w *hdf5Writer) chunkTarget() int { return hdf5ChunkTargetOf(w.opts) }
|
||||
|
||||
// hdf5ChunkTargetOf resolves the chunk target the options ask for.
|
||||
func hdf5ChunkTargetOf(opts HDF5WriteOptions) int {
|
||||
if opts.ChunkBytes == 0 {
|
||||
return hdf5ChunkTarget
|
||||
}
|
||||
return opts.ChunkBytes
|
||||
}
|
||||
|
||||
// hdf5ChunkShape picks a chunk shape that aims at the target bytes: a
|
||||
// dataset that fits the target stays whole, otherwise the last axis is
|
||||
// split first, which keeps the access the files of this package see
|
||||
// (a row of a table, a trace of a signal) inside one chunk.
|
||||
func hdf5ChunkShape(shape []int, width, target int) []int {
|
||||
n := 1
|
||||
for _, d := range shape {
|
||||
n *= d
|
||||
}
|
||||
if n*width <= target {
|
||||
return append([]int{}, shape...)
|
||||
}
|
||||
row := width
|
||||
for _, d := range shape[:len(shape)-1] {
|
||||
row *= d
|
||||
}
|
||||
last := max(target/row, 1)
|
||||
out := append([]int{}, shape[:len(shape)-1]...)
|
||||
return append(out, min(last, shape[len(shape)-1]))
|
||||
}
|
||||
|
||||
// writeChunkTree packs the chunks into version 1 chunk B-tree nodes of
|
||||
// sixty-four children and returns the root's address. An interior
|
||||
// node's key for child i is the first key of child i's subtree, and a
|
||||
// node's level counts its distance from the chunks.
|
||||
func (w *hdf5Writer) writeChunkTree(keys []hdf5ChunkKey, rank, width int) uint64 {
|
||||
kids := make([]hdf5ChunkKid, len(keys))
|
||||
for i, k := range keys {
|
||||
kids[i] = hdf5ChunkKid{first: k, last: k, child: k.addr}
|
||||
}
|
||||
level := byte(0)
|
||||
for len(kids) > 2*hdf5IStoreK {
|
||||
var next []hdf5ChunkKid
|
||||
for start := 0; start < len(kids); start += 2 * hdf5IStoreK {
|
||||
batch := kids[start:min(start+2*hdf5IStoreK, len(kids))]
|
||||
addr := w.writeChunkNode(level, batch, rank, width)
|
||||
next = append(next, hdf5ChunkKid{first: batch[0].first, last: batch[len(batch)-1].last, child: addr})
|
||||
}
|
||||
kids = next
|
||||
level++
|
||||
}
|
||||
return w.writeChunkNode(level, kids, rank, width)
|
||||
}
|
||||
|
||||
// writeChunkNode writes one chunk B-tree node: the first key of every
|
||||
// child, a child pointer each, and the sentinel key that closes the
|
||||
// node. The sibling addresses are the undefined value, as in the group
|
||||
// B-tree.
|
||||
func (w *hdf5Writer) writeChunkNode(level byte, kids []hdf5ChunkKid, rank, width int) uint64 {
|
||||
node := make([]byte, hdf5ChunkNodeSize(rank))
|
||||
copy(node, hdf5Tree)
|
||||
node[4] = 1 // a chunk B-tree
|
||||
node[5] = level
|
||||
binary.LittleEndian.PutUint16(node[6:], uint16(len(kids)))
|
||||
binary.LittleEndian.PutUint64(node[8:], math.MaxUint64) // left sibling
|
||||
binary.LittleEndian.PutUint64(node[16:], math.MaxUint64) // right sibling
|
||||
p := 24
|
||||
writeKey := func(k hdf5ChunkKey) {
|
||||
binary.LittleEndian.PutUint32(node[p:], k.size)
|
||||
binary.LittleEndian.PutUint32(node[p+4:], 0) // the filter mask: every filter ran
|
||||
for i, c := range k.coords {
|
||||
binary.LittleEndian.PutUint64(node[p+8+8*i:], c)
|
||||
}
|
||||
p += 8 + 8*len(k.coords)
|
||||
}
|
||||
for _, k := range kids {
|
||||
writeKey(k.first)
|
||||
binary.LittleEndian.PutUint64(node[p:], k.child)
|
||||
p += 8
|
||||
}
|
||||
writeKey(hdf5ChunkSentinel(kids[len(kids)-1].last, width))
|
||||
return w.bytes(node)
|
||||
}
|
||||
|
||||
// filterEntries lists the configured pipeline in application order for
|
||||
// the filter pipeline message: shuffle before deflate, the order that
|
||||
// lays each element's high-order bytes together for the compressor.
|
||||
// The reference stores the element width as the shuffle's one client
|
||||
// value and the level as deflate's.
|
||||
func (w *hdf5Writer) filterEntries(width int) []hdf5OutFilter {
|
||||
var out []hdf5OutFilter
|
||||
if w.opts.Shuffle {
|
||||
out = append(out, hdf5OutFilter{id: hdf5FilterShuffle, name: "shuffle", values: []uint32{uint32(width)}})
|
||||
}
|
||||
if w.opts.Gzip != 0 {
|
||||
level := w.opts.Gzip
|
||||
if level == -1 {
|
||||
level = 6
|
||||
}
|
||||
out = append(out, hdf5OutFilter{id: hdf5FilterDeflate, name: "deflate", values: []uint32{uint32(level)}})
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
// hdf5ShuffleInto transposes the element bytes of data into out, which
|
||||
// must hold the same length: the first block takes every element's
|
||||
// first byte, the next every second byte. Every byte of out is written.
|
||||
func hdf5ShuffleInto(out, data []byte, width int) {
|
||||
n := len(data) / width
|
||||
for i := range width {
|
||||
for j := range n {
|
||||
out[i*n+j] = data[j*width+i]
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// deflateChunk compresses one staged chunk into the zlib stream the
|
||||
// deflate filter stores: a zlib header in front of the raw deflate
|
||||
// stream, which is what the reader's inflate expects. The compressor
|
||||
// and its output buffer are reused across the chunks of a write;
|
||||
// Reset leaves the compressor in the state a fresh writer holds, so
|
||||
// the stream carries the same bytes a per-chunk writer produced. The
|
||||
// returned bytes are the buffer's, valid until the next call.
|
||||
func (w *hdf5Writer) deflateChunk(data []byte, level int) ([]byte, error) {
|
||||
if w.comp == nil || w.compLevel != level {
|
||||
zw, err := zlib.NewWriterLevel(&w.compBuf, level)
|
||||
if err != nil {
|
||||
return nil, base.Errf("the deflate level %d is refused: %w", level, err)
|
||||
}
|
||||
w.comp, w.compLevel = zw, level
|
||||
} else {
|
||||
w.comp.Reset(&w.compBuf)
|
||||
}
|
||||
w.compBuf.Reset()
|
||||
if _, err := w.comp.Write(data); err != nil {
|
||||
return nil, base.Errf("deflate: %w", err)
|
||||
}
|
||||
if err := w.comp.Close(); err != nil {
|
||||
return nil, base.Errf("deflate: %w", err)
|
||||
}
|
||||
return w.compBuf.Bytes(), nil
|
||||
}
|
||||
@@ -0,0 +1,449 @@
|
||||
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
|
||||
// SPDX-License-Identifier: MIT
|
||||
|
||||
package io
|
||||
|
||||
// Object header and message writing for the HDF5 writer: the messages
|
||||
// here are shaped exactly as the reader's decoders in hdf5.go parse
|
||||
// them, which the fixtures under testdata/h5 pin byte for byte where
|
||||
// it matters.
|
||||
|
||||
import (
|
||||
"encoding/binary"
|
||||
"math"
|
||||
"strconv"
|
||||
"strings"
|
||||
|
||||
"sourcedock.dev/petrbalvin/tensor/internal/base"
|
||||
)
|
||||
|
||||
// hdf5OutMsg is one message of an object header: a type, the payload
|
||||
// the format defines for it.
|
||||
type hdf5OutMsg struct {
|
||||
typ uint16
|
||||
data []byte
|
||||
}
|
||||
|
||||
// hdf5HeaderV1 builds a version 1 object header image, the one the
|
||||
// classic layout uses: the fixed head names the message count and the
|
||||
// size of the message region, and every message is preceded by an
|
||||
// eight-byte header. The message's stored size includes the padding
|
||||
// that keeps the next message on an eight-byte boundary of the header,
|
||||
// which is how the reference library writes every message class.
|
||||
func hdf5HeaderV1(msgs []hdf5OutMsg) ([]byte, error) {
|
||||
if len(msgs) > 0xffff {
|
||||
return nil, base.Errf("SaveHDF5: an object header would carry %d messages, past the %d the format counts", len(msgs), 0xffff)
|
||||
}
|
||||
body := []byte{}
|
||||
for _, m := range msgs {
|
||||
data := hdf5PadField(m.data)
|
||||
if len(data) > 0xffff {
|
||||
return nil, base.Errf("SaveHDF5: a message of %d bytes is past the %d a version 1 header stores", len(data), 0xffff)
|
||||
}
|
||||
body = binary.LittleEndian.AppendUint16(body, m.typ)
|
||||
body = binary.LittleEndian.AppendUint16(body, uint16(len(data)))
|
||||
body = append(body, 0, 0, 0, 0) // message flags and reserved
|
||||
body = append(body, data...)
|
||||
}
|
||||
head := []byte{1, 0}
|
||||
head = binary.LittleEndian.AppendUint16(head, uint16(len(msgs)))
|
||||
head = binary.LittleEndian.AppendUint32(head, 1) // reference count
|
||||
head = binary.LittleEndian.AppendUint32(head, uint32(len(body)))
|
||||
head = binary.LittleEndian.AppendUint32(head, 0) // padding to eight
|
||||
return append(head, body...), nil
|
||||
}
|
||||
|
||||
// hdf5HeaderV2 builds a version 2 object header image, the one the
|
||||
// latest layout uses: the messages are packed without alignment and
|
||||
// the whole header, signature to last message, closes with a lookup3
|
||||
// checksum. The size field is one, two or four bytes, whichever holds
|
||||
// it.
|
||||
func hdf5HeaderV2(msgs []hdf5OutMsg) ([]byte, error) {
|
||||
body := []byte{}
|
||||
for _, m := range msgs {
|
||||
if len(m.data) > 0xffff {
|
||||
return nil, base.Errf("SaveHDF5: a message of %d bytes is past the %d a version 2 header stores", len(m.data), 0xffff)
|
||||
}
|
||||
body = append(body, byte(m.typ))
|
||||
body = binary.LittleEndian.AppendUint16(body, uint16(len(m.data)))
|
||||
body = append(body, 0) // message flags
|
||||
body = append(body, m.data...)
|
||||
}
|
||||
// The size of the size field is itself a mask: the stored value is
|
||||
// the base-2 logarithm of the width, 0 for one byte, 1 for two and
|
||||
// 2 for four, which the reader decodes as 1 << stored.
|
||||
if len(body) > math.MaxUint32 {
|
||||
return nil, base.Errf("SaveHDF5: an object header of %d bytes is past the %d the version 2 size field holds", len(body), uint64(math.MaxUint32))
|
||||
}
|
||||
var logWidth byte
|
||||
switch {
|
||||
case len(body) > 0xffff:
|
||||
logWidth = 2
|
||||
case len(body) > 0xff:
|
||||
logWidth = 1
|
||||
}
|
||||
width := byte(1) << logWidth
|
||||
out := append([]byte{}, hdf5ObjHdr2...)
|
||||
out = append(out, 2, logWidth)
|
||||
for i := range width {
|
||||
out = append(out, byte(len(body)>>(8*i)))
|
||||
}
|
||||
out = append(out, body...)
|
||||
return binary.LittleEndian.AppendUint32(out, hdf5Lookup3(out)), nil
|
||||
}
|
||||
|
||||
// headerAt writes a built header image over a placeholder allocated
|
||||
// earlier, keeping every address that already points here valid.
|
||||
func (w *hdf5Writer) headerAt(addr uint64, image []byte) {
|
||||
copy(w.buf[addr:], image)
|
||||
}
|
||||
|
||||
// hdf5AppendAlign appends the zero bytes that align b to its next
|
||||
// eight-byte boundary.
|
||||
func hdf5AppendAlign(b []byte) []byte {
|
||||
if r := len(b) % 8; r != 0 {
|
||||
return append(b, make([]byte, 8-r)...)
|
||||
}
|
||||
return b
|
||||
}
|
||||
|
||||
// hdf5PadField pads a name or a value field of a message to its next
|
||||
// eight-byte boundary; a field already aligned stays as it is.
|
||||
func hdf5PadField(b []byte) []byte { return hdf5AppendAlign(b) }
|
||||
|
||||
// Datatype messages, version 1. The class bit field's first byte
|
||||
// carries the byte order (cleared: little-endian) and, for fixed
|
||||
// point, the signed bit; the floating-point classes carry the IEEE 754
|
||||
// layout the fixtures store, down to the exponent bias.
|
||||
var (
|
||||
hdf5Float32Type = []byte{
|
||||
0x11, // version 1, class 1 (floating-point)
|
||||
0x20, 0x1f, 0x00, // little-endian, the sign at bit 31
|
||||
0x04, 0x00, 0x00, 0x00, // four-byte elements
|
||||
0x00, 0x00, // bit offset 0
|
||||
0x20, 0x00, // 32 bits of precision
|
||||
0x17, 0x08, 0x00, 0x17, // exponent at 23 of 8, mantissa at 0 of 23
|
||||
0x7f, 0x00, 0x00, 0x00, // exponent bias 127
|
||||
}
|
||||
hdf5Float64Type = []byte{
|
||||
0x11, // version 1, class 1 (floating-point)
|
||||
0x20, 0x3f, 0x00, // little-endian, the sign at bit 63
|
||||
0x08, 0x00, 0x00, 0x00, // eight-byte elements
|
||||
0x00, 0x00, // bit offset 0
|
||||
0x40, 0x00, // 64 bits of precision
|
||||
0x34, 0x0b, 0x00, 0x34, // exponent at 52 of 11, mantissa at 0 of 52
|
||||
0xff, 0x03, 0x00, 0x00, // exponent bias 1023
|
||||
}
|
||||
)
|
||||
|
||||
// hdf5IntType writes a fixed-point datatype message of the given
|
||||
// element width and signedness. The class bit field's byte order bit
|
||||
// stays clear (little-endian) and bit 0x08 carries two's-complement
|
||||
// signedness, the bit decodeType keys its landing on. The message is
|
||||
// twelve bytes: the eight-byte header plus the bit offset and bit
|
||||
// precision the format's fixed-point property table defines.
|
||||
func hdf5IntType(width int, signed bool) []byte {
|
||||
flags := byte(0) // little-endian, unsigned
|
||||
if signed {
|
||||
flags = 0x08 // bit 3: two's complement
|
||||
}
|
||||
b := []byte{0x10, flags, 0, 0} // version 1, class 0 (fixed-point)
|
||||
b = binary.LittleEndian.AppendUint32(b, uint32(width))
|
||||
b = binary.LittleEndian.AppendUint16(b, 0)
|
||||
b = binary.LittleEndian.AppendUint16(b, uint16(8*width))
|
||||
return b
|
||||
}
|
||||
|
||||
// hdf5BoolType writes the boolean enumeration datatype message, in the
|
||||
// exact shape the reader's hdf5EnumBool admits: class 8 with a member
|
||||
// count of two in the class bit field's low sixteen bits and the
|
||||
// reserved byte zero, a base type that is a complete one-byte unsigned
|
||||
// little-endian fixed-point message, the member names each NUL
|
||||
// terminated and padded from its own field start to a multiple of eight
|
||||
// bytes (the message version 1 convention) and the packed member values
|
||||
// 0 and 1 behind the names. The member names carry no semantics; the
|
||||
// values are what the landing reads, and the payload stores them, one
|
||||
// byte per element.
|
||||
func hdf5BoolType() []byte {
|
||||
b := []byte{0x18, 0x02, 0x00, 0x00, 0x01, 0x00, 0x00, 0x00} // version 1, class 8, two members, one-byte values
|
||||
b = append(b, hdf5IntType(1, false)...) // the base type
|
||||
for _, name := range []string{"FALSE", "TRUE"} {
|
||||
start := len(b)
|
||||
b = append(b, name...)
|
||||
b = append(b, 0)
|
||||
for (len(b)-start)%8 != 0 {
|
||||
b = append(b, 0)
|
||||
}
|
||||
}
|
||||
return append(b, 0, 1) // FALSE = 0, TRUE = 1
|
||||
}
|
||||
|
||||
// hdf5StringType writes a fixed-length string datatype message padded
|
||||
// with NUL bytes, the padding the reference writes for fixed strings.
|
||||
func hdf5StringType(width int) []byte {
|
||||
b := []byte{0x13, 0x01, 0, 0} // version 1, class 3, NUL-padded, ASCII
|
||||
return binary.LittleEndian.AppendUint32(b, uint32(width))
|
||||
}
|
||||
|
||||
// hdf5SpaceV1 writes a version 1 dataspace message: the shape, with
|
||||
// the maximum dimensions (the same extents) behind it, as the
|
||||
// reference writes. A scalar dataset declares rank 0.
|
||||
func hdf5SpaceV1(shape []int) []byte {
|
||||
if len(shape) == 0 {
|
||||
return []byte{1, 0, 0, 0, 0, 0, 0, 0}
|
||||
}
|
||||
b := []byte{1, byte(len(shape)), 0x01, 0}
|
||||
b = binary.LittleEndian.AppendUint32(b, 0)
|
||||
for _, d := range shape {
|
||||
b = binary.LittleEndian.AppendUint64(b, uint64(d))
|
||||
}
|
||||
for _, d := range shape {
|
||||
b = binary.LittleEndian.AppendUint64(b, uint64(d))
|
||||
}
|
||||
return b
|
||||
}
|
||||
|
||||
// hdf5SpaceV2 writes a version 2 dataspace message, the one the latest
|
||||
// layout uses: the fourth byte names the dataspace class, which the
|
||||
// reference library sets to simple (one) for every ranked extent.
|
||||
func hdf5SpaceV2(shape []int) []byte {
|
||||
if len(shape) == 0 {
|
||||
return []byte{2, 0, 0, 0} // a scalar: class 0, no dimensions
|
||||
}
|
||||
b := []byte{2, byte(len(shape)), 0x01, 1}
|
||||
for _, d := range shape {
|
||||
b = binary.LittleEndian.AppendUint64(b, uint64(d))
|
||||
}
|
||||
for _, d := range shape {
|
||||
b = binary.LittleEndian.AppendUint64(b, uint64(d))
|
||||
}
|
||||
return b
|
||||
}
|
||||
|
||||
// hdf5FillValueMsg writes the fill value message: the version 2 form
|
||||
// of the classic layout declares a defined, all-zero fill (allocated
|
||||
// incrementally for contiguous storage and late for chunked, as the
|
||||
// reference does), the version 3 form of the latest layout carries no
|
||||
// fill value at all.
|
||||
func hdf5FillValueMsg(latest bool, chunked bool) []byte {
|
||||
if latest {
|
||||
return []byte{3, 0x0a}
|
||||
}
|
||||
if chunked {
|
||||
return []byte{2, 3, 2, 1, 0, 0, 0, 0}
|
||||
}
|
||||
return []byte{2, 2, 2, 1, 0, 0, 0, 0}
|
||||
}
|
||||
|
||||
// hdf5LayoutContiguous writes a version 3 contiguous layout message:
|
||||
// the data's address and its byte extent. An empty dataset names no
|
||||
// storage: the undefined address stands for it.
|
||||
func hdf5LayoutContiguous(addr, size uint64) []byte {
|
||||
b := []byte{3, 1}
|
||||
b = binary.LittleEndian.AppendUint64(b, addr)
|
||||
return binary.LittleEndian.AppendUint64(b, size)
|
||||
}
|
||||
|
||||
// hdf5LayoutChunked writes a version 3 chunked layout message: the
|
||||
// B-tree address, then one more dimension than the dataset has, the
|
||||
// chunk's shape followed by the element size.
|
||||
func hdf5LayoutChunked(addr uint64, chunk []int, width int) []byte {
|
||||
b := []byte{3, 2, byte(len(chunk) + 1)}
|
||||
b = binary.LittleEndian.AppendUint64(b, addr)
|
||||
for _, c := range chunk {
|
||||
b = binary.LittleEndian.AppendUint32(b, uint32(c))
|
||||
}
|
||||
return binary.LittleEndian.AppendUint32(b, uint32(width))
|
||||
}
|
||||
|
||||
// hdf5OutFilter is one entry of a filter pipeline message: the
|
||||
// identifier, the name the reference stores and the client values.
|
||||
type hdf5OutFilter struct {
|
||||
id uint16
|
||||
name string
|
||||
values []uint32
|
||||
}
|
||||
|
||||
// hdf5FilterMessage writes a version 1 filter pipeline message. Every
|
||||
// entry carries the optional flag the reference sets, its name padded
|
||||
// to eight bytes and its client values padded to the same boundary.
|
||||
func hdf5FilterMessage(filters []hdf5OutFilter) []byte {
|
||||
b := []byte{1, byte(len(filters)), 0, 0, 0, 0, 0, 0}
|
||||
for _, f := range filters {
|
||||
name := append([]byte(f.name), 0)
|
||||
b = binary.LittleEndian.AppendUint16(b, f.id)
|
||||
b = binary.LittleEndian.AppendUint16(b, uint16(len(name)))
|
||||
b = binary.LittleEndian.AppendUint16(b, 1)
|
||||
b = binary.LittleEndian.AppendUint16(b, uint16(len(f.values)))
|
||||
b = hdf5PadField(append(b, name...))
|
||||
for _, v := range f.values {
|
||||
b = binary.LittleEndian.AppendUint32(b, v)
|
||||
}
|
||||
b = hdf5AppendAlign(b)
|
||||
}
|
||||
return b
|
||||
}
|
||||
|
||||
// hdf5AttrMessage writes one version 1 attribute message: the name
|
||||
// padded to eight bytes, the datatype padded to eight, the dataspace,
|
||||
// then the value parsed out of its text. The reader renders every
|
||||
// attribute it reads as text, so the writer parses text back into the
|
||||
// typed attribute it names: a whole number becomes int64, a decimal
|
||||
// float64, a bracketed list an int64 or float64 array and anything
|
||||
// else a fixed-length string.
|
||||
func hdf5AttrMessage(a hdf5Attr) (hdf5OutMsg, error) {
|
||||
dtype, value, shape, err := hdf5AttrValue(a.name, a.text)
|
||||
if err != nil {
|
||||
return hdf5OutMsg{}, err
|
||||
}
|
||||
// The name size counts the name and its NUL terminator; the field
|
||||
// itself is padded to eight bytes.
|
||||
nameField := hdf5PadField(append([]byte(a.name), 0))
|
||||
space := hdf5SpaceV1(shape)
|
||||
body := []byte{1, 0}
|
||||
body = binary.LittleEndian.AppendUint16(body, uint16(len(a.name)+1))
|
||||
body = binary.LittleEndian.AppendUint16(body, uint16(len(dtype)))
|
||||
body = binary.LittleEndian.AppendUint16(body, uint16(len(space)))
|
||||
body = append(body, nameField...)
|
||||
body = hdf5PadField(append(body, dtype...))
|
||||
body = append(body, space...)
|
||||
body = append(body, value...)
|
||||
return hdf5OutMsg{typ: hdf5MsgAttribute, data: body}, nil
|
||||
}
|
||||
|
||||
// hdf5AttrValue parses an attribute's text into a datatype message,
|
||||
// the raw value bytes and the dataspace shape (nil for a scalar). It
|
||||
// is the inverse of the reader's attribute rendering: FormatInt and
|
||||
// FormatFloat output parse back to the values they were printed from,
|
||||
// bit for bit, and every other text becomes a fixed-length string.
|
||||
func hdf5AttrValue(name, text string) ([]byte, []byte, []int, error) {
|
||||
if strings.HasPrefix(text, "[") && strings.HasSuffix(text, "]") {
|
||||
inner := strings.TrimSpace(text[1 : len(text)-1])
|
||||
if inner == "" {
|
||||
// An empty array: an int64 attribute of extent zero, which
|
||||
// the reader renders back as "[]".
|
||||
return hdf5IntType(8, true), nil, []int{0}, nil
|
||||
}
|
||||
parts := strings.Split(inner, ", ")
|
||||
ints := make([]int64, 0, len(parts))
|
||||
floats := make([]float64, 0, len(parts))
|
||||
asFloat := false
|
||||
for i, p := range parts {
|
||||
if v, err := strconv.ParseInt(p, 10, 64); err == nil && !asFloat {
|
||||
ints = append(ints, v)
|
||||
floats = append(floats, float64(v))
|
||||
continue
|
||||
}
|
||||
v, err := strconv.ParseFloat(p, 64)
|
||||
if err != nil {
|
||||
return nil, nil, nil, base.Errf("the attribute %q holds the array value %q whose element %d is not a number", name, text, i)
|
||||
}
|
||||
asFloat = true
|
||||
floats = append(floats, v)
|
||||
}
|
||||
if asFloat {
|
||||
raw := make([]byte, 0, 8*len(floats))
|
||||
for _, v := range floats {
|
||||
raw = binary.LittleEndian.AppendUint64(raw, math.Float64bits(v))
|
||||
}
|
||||
return hdf5Float64Type, raw, []int{len(floats)}, nil
|
||||
}
|
||||
raw := make([]byte, 0, 8*len(ints))
|
||||
for _, v := range ints {
|
||||
raw = binary.LittleEndian.AppendUint64(raw, uint64(v))
|
||||
}
|
||||
return hdf5IntType(8, true), raw, []int{len(ints)}, nil
|
||||
}
|
||||
if v, err := strconv.ParseInt(text, 10, 64); err == nil {
|
||||
return hdf5IntType(8, true), binary.LittleEndian.AppendUint64(nil, uint64(v)), nil, nil
|
||||
}
|
||||
if v, err := strconv.ParseFloat(text, 64); err == nil {
|
||||
return hdf5Float64Type, binary.LittleEndian.AppendUint64(nil, math.Float64bits(v)), nil, nil
|
||||
}
|
||||
width := max(len(text), 1)
|
||||
value := append([]byte(text), make([]byte, width-len(text))...)
|
||||
return hdf5StringType(width), value, nil, nil
|
||||
}
|
||||
|
||||
// writeDataset writes one dataset: the dataspace, datatype and fill
|
||||
// value messages, the filter pipeline and layout of its storage and
|
||||
// its attributes, then the header of the layout the file uses.
|
||||
func (w *hdf5Writer) writeDataset(s *hdf5OutSet) (uint64, error) {
|
||||
msgs := make([]hdf5OutMsg, 0, 6)
|
||||
if w.latest {
|
||||
msgs = append(msgs, hdf5OutMsg{typ: hdf5MsgDataspace, data: hdf5SpaceV2(s.shape)})
|
||||
} else {
|
||||
msgs = append(msgs, hdf5OutMsg{typ: hdf5MsgDataspace, data: hdf5SpaceV1(s.shape)})
|
||||
}
|
||||
switch s.class {
|
||||
case 0:
|
||||
msgs = append(msgs, hdf5OutMsg{typ: hdf5MsgDatatype, data: hdf5IntType(s.width, s.signed)})
|
||||
case 1:
|
||||
if s.width == 4 {
|
||||
msgs = append(msgs, hdf5OutMsg{typ: hdf5MsgDatatype, data: hdf5Float32Type})
|
||||
} else {
|
||||
msgs = append(msgs, hdf5OutMsg{typ: hdf5MsgDatatype, data: hdf5Float64Type})
|
||||
}
|
||||
case 3:
|
||||
msgs = append(msgs, hdf5OutMsg{typ: hdf5MsgDatatype, data: hdf5StringType(s.width)})
|
||||
case 8:
|
||||
msgs = append(msgs, hdf5OutMsg{typ: hdf5MsgDatatype, data: hdf5BoolType()})
|
||||
default:
|
||||
// Every class the plan can produce has its message case above;
|
||||
// an unlisted one would emit a file with no datatype message,
|
||||
// which the reader refuses later with less context than this.
|
||||
return 0, base.Errf("SaveHDF5: dataset %q: the writer emits no datatype message for class %d", s.path, s.class)
|
||||
}
|
||||
// Chunked storage is the filter pipeline's only carrier, so a
|
||||
// dataset with filters goes chunked; strings keep their raw bytes,
|
||||
// and a dataset with no elements has nothing for the filters to
|
||||
// compress, so it stays contiguous at an undefined address either
|
||||
// way.
|
||||
chunked := len(s.shape) > 0 && s.nbytes > 0 && (w.opts.Gzip != 0 || w.opts.Shuffle) && s.class != 3
|
||||
msgs = append(msgs, hdf5OutMsg{typ: hdf5MsgFillValue, data: hdf5FillValueMsg(w.latest, chunked)})
|
||||
if chunked {
|
||||
entries := w.filterEntries(s.width)
|
||||
if len(entries) > 0 {
|
||||
msgs = append(msgs, hdf5OutMsg{typ: hdf5MsgFilterPipeline, data: hdf5FilterMessage(entries)})
|
||||
}
|
||||
tree, chunk, err := w.writeChunks(s)
|
||||
if err != nil {
|
||||
return 0, base.Errf("SaveHDF5: dataset %q: %w", s.path, err)
|
||||
}
|
||||
msgs = append(msgs, hdf5OutMsg{typ: hdf5MsgDataLayout, data: hdf5LayoutChunked(tree, chunk, s.width)})
|
||||
} else {
|
||||
addr := uint64(math.MaxUint64)
|
||||
var size uint64
|
||||
if s.nbytes > 0 {
|
||||
// The payload is encoded where the format will read it:
|
||||
// the reservation is its final address, so the values pass
|
||||
// through the writer once instead of building a block and
|
||||
// then copying it in.
|
||||
addr = w.reserve(s.nbytes)
|
||||
if encErr := s.encode(w.buf[addr:]); encErr != nil {
|
||||
return 0, base.Errf("SaveHDF5: dataset %q: %w", s.path, encErr)
|
||||
}
|
||||
w.pad8()
|
||||
size = uint64(s.nbytes)
|
||||
}
|
||||
msgs = append(msgs, hdf5OutMsg{typ: hdf5MsgDataLayout, data: hdf5LayoutContiguous(addr, size)})
|
||||
}
|
||||
for _, a := range s.attrs {
|
||||
m, err := hdf5AttrMessage(a)
|
||||
if err != nil {
|
||||
return 0, base.Errf("SaveHDF5: dataset %q: %w", s.path, err)
|
||||
}
|
||||
msgs = append(msgs, m)
|
||||
}
|
||||
var image []byte
|
||||
var err error
|
||||
if w.latest {
|
||||
image, err = hdf5HeaderV2(msgs)
|
||||
} else {
|
||||
image, err = hdf5HeaderV1(msgs)
|
||||
}
|
||||
if err != nil {
|
||||
return 0, base.Errf("SaveHDF5: dataset %q: %w", s.path, err)
|
||||
}
|
||||
return w.bytes(image), nil
|
||||
}
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,34 @@
|
||||
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
|
||||
// SPDX-License-Identifier: MIT
|
||||
|
||||
package io
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"sourcedock.dev/petrbalvin/tensor/internal/core"
|
||||
)
|
||||
|
||||
// mustFloats builds a float array, failing the test on a bad shape.
|
||||
// Without an explicit shape it defaults to a vector of len(vals).
|
||||
func mustFloats(t *testing.T, vals []float64, shape ...int) *core.Array {
|
||||
t.Helper()
|
||||
if len(shape) == 0 {
|
||||
shape = []int{len(vals)}
|
||||
}
|
||||
a, err := core.FromFloats(vals, shape...)
|
||||
if err != nil {
|
||||
t.Fatalf("FromFloats(%v, %v): %v", vals, shape, err)
|
||||
}
|
||||
return a
|
||||
}
|
||||
|
||||
// mustComplexes builds a complex array, failing the test on a bad shape.
|
||||
func mustComplexes(t *testing.T, vals []complex128, shape ...int) *core.Array {
|
||||
t.Helper()
|
||||
a, err := core.FromComplexes(vals, shape...)
|
||||
if err != nil {
|
||||
t.Fatalf("FromComplexes(%v, %v): %v", vals, shape, err)
|
||||
}
|
||||
return a
|
||||
}
|
||||
@@ -0,0 +1,267 @@
|
||||
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
|
||||
// SPDX-License-Identifier: MIT
|
||||
|
||||
package io
|
||||
|
||||
import (
|
||||
"encoding/binary"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"testing"
|
||||
)
|
||||
|
||||
// Hostile inputs: a data file is untrusted input, so every size field
|
||||
// the header carries must be checked against the bytes actually
|
||||
// present before it is used. These cases pin the contract that a
|
||||
// malformed file is an error, never a panic and never an allocation
|
||||
// sized by the header alone.
|
||||
|
||||
// hostileNCName writes a 4-byte-padded NetCDF name.
|
||||
func hostileNCName(b []byte, name string) []byte {
|
||||
b = binary.BigEndian.AppendUint32(b, uint32(len(name)))
|
||||
b = append(b, name...)
|
||||
if pad := (4 - len(name)%4) % 4; pad != 0 {
|
||||
b = append(b, make([]byte, pad)...)
|
||||
}
|
||||
return b
|
||||
}
|
||||
|
||||
// ncHostileHeader builds a minimal CDF-1 file: magic, numrecs, a
|
||||
// dimension list with the given lengths, an absent attribute list, and
|
||||
// one NC_DOUBLE variable over those dimensions. Nothing follows the
|
||||
// header, so any declared data is truncated by construction.
|
||||
func ncHostileHeader(names []string, lengths []uint32) []byte {
|
||||
b := []byte{'C', 'D', 'F', 1}
|
||||
b = binary.BigEndian.AppendUint32(b, 0) // numrecs
|
||||
b = binary.BigEndian.AppendUint32(b, ncTagDimension)
|
||||
b = binary.BigEndian.AppendUint32(b, uint32(len(lengths)))
|
||||
for i, l := range lengths {
|
||||
b = hostileNCName(b, names[i])
|
||||
b = binary.BigEndian.AppendUint32(b, l)
|
||||
}
|
||||
b = binary.BigEndian.AppendUint32(b, 0) // absent attribute list
|
||||
b = binary.BigEndian.AppendUint32(b, 0)
|
||||
b = binary.BigEndian.AppendUint32(b, ncTagVariable)
|
||||
b = binary.BigEndian.AppendUint32(b, 1)
|
||||
b = hostileNCName(b, "v")
|
||||
b = binary.BigEndian.AppendUint32(b, uint32(len(lengths)))
|
||||
for i := range lengths {
|
||||
b = binary.BigEndian.AppendUint32(b, uint32(i))
|
||||
}
|
||||
b = binary.BigEndian.AppendUint32(b, 0) // absent variable attributes
|
||||
b = binary.BigEndian.AppendUint32(b, 0)
|
||||
b = binary.BigEndian.AppendUint32(b, ncTypeDouble)
|
||||
b = binary.BigEndian.AppendUint32(b, 0) // vsize
|
||||
b = binary.BigEndian.AppendUint32(b, 0) // begin
|
||||
return b
|
||||
}
|
||||
|
||||
// writeHostile drops the bytes into a temp file and returns the path.
|
||||
func writeHostile(t *testing.T, name string, data []byte) string {
|
||||
t.Helper()
|
||||
path := filepath.Join(t.TempDir(), name)
|
||||
if err := os.WriteFile(path, data, 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
return path
|
||||
}
|
||||
|
||||
// TestLoadNetCDFHostileCounts covers the three ways a declared count
|
||||
// can be used to demand memory the file cannot back: a product that
|
||||
// wraps to a small positive number, a product that wraps negative, and
|
||||
// record counts that are simply larger than the file.
|
||||
func TestLoadNetCDFHostileCounts(t *testing.T) {
|
||||
cases := []struct {
|
||||
name string
|
||||
names []string
|
||||
lengths []uint32
|
||||
}{
|
||||
{
|
||||
// 2^21 * 2^20 * 2^20 = 2^61, which times the 8-byte
|
||||
// element width wraps to zero in a 64-bit multiply.
|
||||
name: "product wraps to zero",
|
||||
names: []string{"a", "b", "c"},
|
||||
lengths: []uint32{1 << 21, 1 << 20, 1 << 20},
|
||||
},
|
||||
{
|
||||
// 4294967295 squared wraps to a negative product.
|
||||
name: "product wraps negative",
|
||||
names: []string{"a", "b"},
|
||||
lengths: []uint32{4294967295, 4294967295},
|
||||
},
|
||||
{
|
||||
name: "one dimension longer than the file",
|
||||
names: []string{"a"},
|
||||
lengths: []uint32{1 << 30},
|
||||
},
|
||||
}
|
||||
for _, tc := range cases {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
path := writeHostile(t, "hostile.nc", ncHostileHeader(tc.names, tc.lengths))
|
||||
if _, _, _, err := LoadNetCDF(path); err == nil {
|
||||
t.Fatal("expected an error for a header whose data cannot fit the file")
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// TestLoadNetCDFHostileRecordCounts pins the list caps: a header that
|
||||
// declares millions of dimensions, attributes or variables in a file
|
||||
// of a few bytes must fail before the list is allocated.
|
||||
func TestLoadNetCDFHostileRecordCounts(t *testing.T) {
|
||||
const declared = 2000000
|
||||
|
||||
dimCount := []byte{'C', 'D', 'F', 1}
|
||||
dimCount = binary.BigEndian.AppendUint32(dimCount, 0)
|
||||
dimCount = binary.BigEndian.AppendUint32(dimCount, ncTagDimension)
|
||||
dimCount = binary.BigEndian.AppendUint32(dimCount, declared)
|
||||
|
||||
attrCount := []byte{'C', 'D', 'F', 1}
|
||||
attrCount = binary.BigEndian.AppendUint32(attrCount, 0)
|
||||
attrCount = binary.BigEndian.AppendUint32(attrCount, 0) // absent dimensions
|
||||
attrCount = binary.BigEndian.AppendUint32(attrCount, 0)
|
||||
attrCount = binary.BigEndian.AppendUint32(attrCount, ncTagAttribute)
|
||||
attrCount = binary.BigEndian.AppendUint32(attrCount, declared)
|
||||
|
||||
varCount := []byte{'C', 'D', 'F', 1}
|
||||
varCount = binary.BigEndian.AppendUint32(varCount, 0)
|
||||
varCount = binary.BigEndian.AppendUint32(varCount, 0) // absent dimensions
|
||||
varCount = binary.BigEndian.AppendUint32(varCount, 0)
|
||||
varCount = binary.BigEndian.AppendUint32(varCount, 0) // absent attributes
|
||||
varCount = binary.BigEndian.AppendUint32(varCount, 0)
|
||||
varCount = binary.BigEndian.AppendUint32(varCount, ncTagVariable)
|
||||
varCount = binary.BigEndian.AppendUint32(varCount, declared)
|
||||
|
||||
cases := []struct {
|
||||
name string
|
||||
data []byte
|
||||
}{
|
||||
{"dimensions", dimCount},
|
||||
{"attributes", attrCount},
|
||||
{"variables", varCount},
|
||||
}
|
||||
for _, tc := range cases {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
path := writeHostile(t, "counts.nc", tc.data)
|
||||
_, _, _, err := LoadNetCDF(path)
|
||||
if err == nil {
|
||||
t.Fatalf("a %d-byte file declaring %d %s must be refused", len(tc.data), declared, tc.name)
|
||||
}
|
||||
if !strings.Contains(err.Error(), "remaining bytes") {
|
||||
t.Fatalf("error = %v, want the declared-count bound", err)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// TestLoadNetCDFTruncatedVariable pins the data-block check: a header
|
||||
// that points its variable past the end of the file fails without
|
||||
// reading beyond it.
|
||||
func TestLoadNetCDFTruncatedVariable(t *testing.T) {
|
||||
data := ncHostileHeader([]string{"a"}, []uint32{4})
|
||||
binary.BigEndian.PutUint32(data[len(data)-4:], 1<<30) // begin past EOF
|
||||
path := writeHostile(t, "short.nc", data)
|
||||
if _, _, _, err := LoadNetCDF(path); err == nil {
|
||||
t.Fatal("expected an error for a variable whose data lies past the end of the file")
|
||||
}
|
||||
}
|
||||
|
||||
// TestLoadFITSTableTruncatedHostile pins the table data guard: a file
|
||||
// that ends inside the header block has no data at all, whether or not
|
||||
// the column forms give it a row size to divide by.
|
||||
func TestLoadFITSTableTruncatedHostile(t *testing.T) {
|
||||
// No data block follows the header: the cards stop at 640 bytes.
|
||||
header := func(form string) []byte {
|
||||
var b []byte
|
||||
for _, body := range []string{
|
||||
"XTENSION= 'BINTABLE'", "BITPIX = 8", "NAXIS = 2",
|
||||
"NAXIS1 = 4", "NAXIS2 = 2", "TFIELDS = 1",
|
||||
"TFORM1 = '" + form + strings.Repeat(" ", 8-len(form)) + "'", "END",
|
||||
} {
|
||||
b = append(b, card(body)...)
|
||||
}
|
||||
return b
|
||||
}
|
||||
for _, form := range []string{"X", "D"} {
|
||||
t.Run(form, func(t *testing.T) {
|
||||
data := header(form)
|
||||
if len(data) != 640 {
|
||||
t.Fatalf("the test header is %d bytes, want 640", len(data))
|
||||
}
|
||||
path := writeHostile(t, "trunc.fits", data)
|
||||
if _, err := LoadFITSTable(path); err == nil {
|
||||
t.Fatal("expected an error for a table whose data block is missing")
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// TestLoadNetCDFUnusedLongDimension pins the other side of the bound: a
|
||||
// header may declare a dimension longer than the data any variable
|
||||
// uses, because the format allows it, so the reader must accept the
|
||||
// file rather than refuse a legal one.
|
||||
func TestLoadNetCDFUnusedLongDimension(t *testing.T) {
|
||||
// One dimension of a million elements, one variable using none of
|
||||
// it: the count product is 1, so nothing is read.
|
||||
var b []byte
|
||||
b = append(b, 'C', 'D', 'F', 1)
|
||||
b = binary.BigEndian.AppendUint32(b, 0) // numrecs
|
||||
b = binary.BigEndian.AppendUint32(b, 10) // NC_DIMENSION
|
||||
b = binary.BigEndian.AppendUint32(b, 1)
|
||||
b = hostileNCName(b, "big")
|
||||
b = binary.BigEndian.AppendUint32(b, 1<<20)
|
||||
b = binary.BigEndian.AppendUint32(b, 0) // absent attributes
|
||||
b = binary.BigEndian.AppendUint32(b, 0)
|
||||
b = binary.BigEndian.AppendUint32(b, 11) // NC_VARIABLE
|
||||
b = binary.BigEndian.AppendUint32(b, 1)
|
||||
b = hostileNCName(b, "scalar")
|
||||
b = binary.BigEndian.AppendUint32(b, 0) // rank 0
|
||||
b = binary.BigEndian.AppendUint32(b, 0) // absent attributes
|
||||
b = binary.BigEndian.AppendUint32(b, 0)
|
||||
b = binary.BigEndian.AppendUint32(b, 6) // NC_DOUBLE
|
||||
b = binary.BigEndian.AppendUint32(b, 8) // vsize
|
||||
b = binary.BigEndian.AppendUint32(b, 0) // begin
|
||||
b = append(b, 0, 0, 0, 0, 0, 0, 0, 0) // one double at offset 0
|
||||
path := writeHostile(t, "unused.nc", b)
|
||||
dims, vars, _, err := LoadNetCDF(path)
|
||||
if err != nil {
|
||||
t.Fatalf("LoadNetCDF refused a legal file: %v", err)
|
||||
}
|
||||
if len(dims) != 1 || dims[0].Length != 1<<20 {
|
||||
t.Fatalf("dims = %+v, want one dimension of 2^20", dims)
|
||||
}
|
||||
if len(vars) != 1 || vars[0].Values.Len() != 1 {
|
||||
t.Fatalf("vars = %+v, want one scalar", vars)
|
||||
}
|
||||
}
|
||||
|
||||
// TestLoadHDF5ChunkTreeCycle pins the visited set on the chunk B-tree
|
||||
// walk: an inner node that lists itself among its children must be
|
||||
// refused. Without the set the walk multiplies into entries^depth node
|
||||
// visits before the depth guard can fire, so a few hundred crafted
|
||||
// bytes never terminate.
|
||||
func TestLoadHDF5ChunkTreeCycle(t *testing.T) {
|
||||
const entries = 512
|
||||
rank := 1
|
||||
keySize := 8 + 8*(rank+1)
|
||||
entrySize := keySize + 8
|
||||
buf := make([]byte, 24+entries*entrySize)
|
||||
copy(buf, []byte("TREE"))
|
||||
buf[4] = 1 // a chunk node
|
||||
buf[5] = 1 // inner level: every child is recursed, not placed
|
||||
binary.LittleEndian.PutUint16(buf[6:], entries)
|
||||
for i := range entries {
|
||||
p := 24 + i*entrySize
|
||||
binary.LittleEndian.PutUint32(buf[p:], 16) // chunk size
|
||||
binary.LittleEndian.PutUint32(buf[p+4:], 0) // filter mask
|
||||
// The key offsets stay zero; the child address points at
|
||||
// this very node.
|
||||
binary.LittleEndian.PutUint64(buf[p+keySize:], 0)
|
||||
}
|
||||
f := &hdf5File{data: buf, offSize: 8, lenSize: 8}
|
||||
place := func(uint64, []uint64, uint32, int) error { return nil }
|
||||
if err := f.chunkTree(0, hdf5Layout{}, rank, map[uint64]bool{}, place, 0); err == nil || !strings.Contains(err.Error(), "revisits") {
|
||||
t.Fatalf("chunkTree on a self-referencing node: %v", err)
|
||||
}
|
||||
}
|
||||
+181
@@ -0,0 +1,181 @@
|
||||
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
|
||||
// SPDX-License-Identifier: MIT
|
||||
|
||||
package io
|
||||
|
||||
import (
|
||||
"encoding/binary"
|
||||
"math"
|
||||
"os"
|
||||
"syscall"
|
||||
"unsafe"
|
||||
|
||||
"sourcedock.dev/petrbalvin/tensor/internal/base"
|
||||
"sourcedock.dev/petrbalvin/tensor/internal/core"
|
||||
)
|
||||
|
||||
// Memory-mapped arrays. A file of native-endian numbers can back an
|
||||
// Array directly: the operating system maps the file's pages into the
|
||||
// address space and the array reads them in place, so a data cube far
|
||||
// larger than RAM opens instantly and only the touched pages ever
|
||||
// reach memory. The mapping is read-only, which matches the arrays'
|
||||
// immutability contract exactly, and the caller releases it with the
|
||||
// returned function once the numbers are no longer needed.
|
||||
|
||||
// MapFloats maps n float64 values of path, starting at byte offset,
|
||||
// into a read-only one-dimensional array. The values must have been
|
||||
// written in the machine's native byte order (binary.NativeEndian).
|
||||
// The array is a live view of the mapping: release unmaps it, and any
|
||||
// use of the array afterwards is a use-after-free, so release comes
|
||||
// strictly last. A negative offset, a non-positive count, a missing
|
||||
// file or a file that does not hold all n values is an error.
|
||||
func MapFloats(path string, offset int64, n int) (a *core.Array, release func() error, err error) {
|
||||
const name = "MapFloats"
|
||||
length, err := countBytes(name, n, 8)
|
||||
if err != nil {
|
||||
return nil, nil, err
|
||||
}
|
||||
raw, release, err := mapRegion(path, offset, length, 8, name)
|
||||
if err != nil {
|
||||
return nil, nil, err
|
||||
}
|
||||
// SAFETY: countBytes checked that n is positive and that n*8 does not
|
||||
// overflow, and mapRegion checked that the file holds that many bytes
|
||||
// from an offset that is a multiple of 8 and returned a mapping whose
|
||||
// base is page-aligned with exactly that intra-page offset, so the
|
||||
// pointer is 8-aligned and n float64 values lie inside the mapping.
|
||||
// The mapping stays alive until the caller runs release.
|
||||
values := unsafe.Slice((*float64)(unsafe.Pointer(&raw[0])), n)
|
||||
a, err = core.FromFloatSlice(values, n)
|
||||
if err != nil {
|
||||
_ = release()
|
||||
return nil, nil, base.Errf("%s: %w", name, err)
|
||||
}
|
||||
return a, release, nil
|
||||
}
|
||||
|
||||
// MapFloat32s maps n float32 values of path into a read-only array,
|
||||
// with the same contract as MapFloats.
|
||||
func MapFloat32s(path string, offset int64, n int) (a *core.Array, release func() error, err error) {
|
||||
const name = "MapFloat32s"
|
||||
length, err := countBytes(name, n, 4)
|
||||
if err != nil {
|
||||
return nil, nil, err
|
||||
}
|
||||
raw, release, err := mapRegion(path, offset, length, 4, name)
|
||||
if err != nil {
|
||||
return nil, nil, err
|
||||
}
|
||||
// SAFETY: as in MapFloats, with the float32 alignment of 4.
|
||||
values := unsafe.Slice((*float32)(unsafe.Pointer(&raw[0])), n)
|
||||
a, err = core.FromFloat32Slice(values, n)
|
||||
if err != nil {
|
||||
_ = release()
|
||||
return nil, nil, base.Errf("%s: %w", name, err)
|
||||
}
|
||||
return a, release, nil
|
||||
}
|
||||
|
||||
// MapInts maps n int64 values of path into a read-only array, with
|
||||
// the same contract as MapFloats.
|
||||
func MapInts(path string, offset int64, n int) (a *core.Array, release func() error, err error) {
|
||||
const name = "MapInts"
|
||||
length, err := countBytes(name, n, 8)
|
||||
if err != nil {
|
||||
return nil, nil, err
|
||||
}
|
||||
raw, release, err := mapRegion(path, offset, length, 8, name)
|
||||
if err != nil {
|
||||
return nil, nil, err
|
||||
}
|
||||
// SAFETY: as in MapFloats; an int64 has the same size and alignment.
|
||||
values := unsafe.Slice((*int64)(unsafe.Pointer(&raw[0])), n)
|
||||
a, err = core.IntsFromArray(values, n)
|
||||
if err != nil {
|
||||
_ = release()
|
||||
return nil, nil, base.Errf("%s: %w", name, err)
|
||||
}
|
||||
return a, release, nil
|
||||
}
|
||||
|
||||
// countBytes converts an element count into a byte length, refusing the
|
||||
// counts whose byte size does not fit an int64. Multiplying first and
|
||||
// checking the product afterwards is the failure the review found: for
|
||||
// n = 2^61+3 float64 values the product wraps to 24, which passes every
|
||||
// file-size check, and the typed view is then built with the original n,
|
||||
// which unsafe.Slice rejects with a panic instead of an error.
|
||||
func countBytes(name string, n int, width int64) (int64, error) {
|
||||
if n <= 0 {
|
||||
return 0, base.Errf("%s: the element count must be positive, got %d", name, n)
|
||||
}
|
||||
if int64(n) > math.MaxInt64/width {
|
||||
return 0, base.Errf("%s: %d elements of %d bytes take more bytes than a length can address",
|
||||
name, n, width)
|
||||
}
|
||||
return int64(n) * width, nil
|
||||
}
|
||||
|
||||
// mapRegion maps length bytes of path at offset read-only and returns
|
||||
// the byte slice viewing them plus a release that unmaps. The count,
|
||||
// the offset and the element alignment are checked against the file's
|
||||
// size before the mapping, so a short file or a misaligned request
|
||||
// fails with an error instead of a fault or a misaligned array.
|
||||
func mapRegion(path string, offset, length, align int64, name string) ([]byte, func() error, error) {
|
||||
if offset < 0 {
|
||||
return nil, nil, base.Errf("%s: offset must not be negative, got %d", name, offset)
|
||||
}
|
||||
// The caller turns the bytes into a typed slice with unsafe, which
|
||||
// requires the address to satisfy the element's alignment: an
|
||||
// offset that is not a multiple of the element size would hand back
|
||||
// a misaligned array.
|
||||
if offset%align != 0 {
|
||||
return nil, nil, base.Errf("%s: offset %d is not a multiple of %d, the element size", name, offset, align)
|
||||
}
|
||||
if length <= 0 {
|
||||
return nil, nil, base.Errf("%s: length must be positive, got %d", name, length)
|
||||
}
|
||||
file, err := os.Open(path)
|
||||
if err != nil {
|
||||
return nil, nil, base.Errf("%s: %w", name, err)
|
||||
}
|
||||
defer file.Close()
|
||||
info, err := file.Stat()
|
||||
if err != nil {
|
||||
return nil, nil, base.Errf("%s: %w", name, err)
|
||||
}
|
||||
if info.IsDir() {
|
||||
return nil, nil, base.Errf("%s: %s is a directory", name, path)
|
||||
}
|
||||
if size := info.Size(); offset > size || length > size-offset {
|
||||
return nil, nil, base.Errf("%s: %s holds %d bytes, %d needed from offset %d",
|
||||
name, path, size, length, offset)
|
||||
}
|
||||
// The mapping call demands a page-aligned offset; misaligned
|
||||
// requests map from the page boundary below and the returned slice
|
||||
// skips the intra-page part.
|
||||
page := int64(os.Getpagesize())
|
||||
pageBase := offset / page * page
|
||||
raw, err := syscall.Mmap(int(file.Fd()), pageBase, int(length+offset-pageBase), syscall.PROT_READ, syscall.MAP_SHARED)
|
||||
if err != nil {
|
||||
return nil, nil, base.Errf("%s: mapping %s failed: %w", name, path, err)
|
||||
}
|
||||
view := raw[offset-pageBase:]
|
||||
return view, func() error {
|
||||
if err := syscall.Munmap(raw); err != nil {
|
||||
return base.Errf("%s: unmapping %s failed: %w", name, path, err)
|
||||
}
|
||||
return nil
|
||||
}, nil
|
||||
}
|
||||
|
||||
// SaveNativeFloats writes floats to path in the machine's native byte
|
||||
// order, the format MapFloats reads back. It exists so a mapping test
|
||||
// or tool can produce its own data without reaching for encoding
|
||||
// details; a zero offset aligns it with MapFloats' contract.
|
||||
func SaveNativeFloats(path string, values []float64) error {
|
||||
buf := make([]byte, 8*len(values))
|
||||
for i, v := range values {
|
||||
binary.NativeEndian.PutUint64(buf[i*8:], math.Float64bits(v))
|
||||
}
|
||||
return os.WriteFile(path, buf, 0o644)
|
||||
}
|
||||
+165
@@ -0,0 +1,165 @@
|
||||
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
|
||||
// SPDX-License-Identifier: MIT
|
||||
|
||||
package io
|
||||
|
||||
import (
|
||||
"encoding/binary"
|
||||
"math"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"sourcedock.dev/petrbalvin/tensor/internal/core"
|
||||
"testing"
|
||||
)
|
||||
|
||||
// TestMapFloatsRoundTrip writes floats natively, maps them back, and
|
||||
// checks spot values far apart, which is exactly the sparse-touch
|
||||
// pattern big-file mapping exists for.
|
||||
func TestMapFloatsRoundTrip(t *testing.T) {
|
||||
const n = 4096
|
||||
path := filepath.Join(t.TempDir(), "values.bin")
|
||||
values := make([]float64, n)
|
||||
for i := range n {
|
||||
values[i] = math.Sin(float64(i)) * 1e6
|
||||
}
|
||||
if err := SaveNativeFloats(path, values); err != nil {
|
||||
t.Fatalf("SaveNativeFloats: %v", err)
|
||||
}
|
||||
a, release, err := MapFloats(path, 0, n)
|
||||
if err != nil {
|
||||
t.Fatalf("MapFloats: %v", err)
|
||||
}
|
||||
if a.Len() != n || a.Dtype() != core.Float {
|
||||
t.Fatalf("mapped array %s of length %d", a.Dtype(), a.Len())
|
||||
}
|
||||
for _, i := range []int{0, 1, 999, 2048, n - 1} {
|
||||
if a.FloatAt(i) != values[i] {
|
||||
t.Fatalf("value %d = %.14g, want %.14g", i, a.FloatAt(i), values[i])
|
||||
}
|
||||
}
|
||||
// A reshape keeps the view: strides read straight from the mapping.
|
||||
square, err := core.Reshape(a, 64, 64)
|
||||
if err != nil {
|
||||
t.Fatalf("Reshape: %v", err)
|
||||
}
|
||||
if square.FloatAt(3*64+7) != values[3*64+7] {
|
||||
t.Fatal("the reshaped view lost the mapping")
|
||||
}
|
||||
if err := release(); err != nil {
|
||||
t.Fatalf("release: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
// TestMapFloat32sAndInts covers the other two element types.
|
||||
func TestMapFloat32sAndInts(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
const n = 100
|
||||
f32path := filepath.Join(dir, "f32.bin")
|
||||
var buf []byte
|
||||
var word [4]byte
|
||||
want32 := make([]float32, n)
|
||||
for i := range n {
|
||||
want32[i] = float32(i) / 7
|
||||
binary.NativeEndian.PutUint32(word[:], math.Float32bits(want32[i]))
|
||||
buf = append(buf, word[:]...)
|
||||
}
|
||||
if err := os.WriteFile(f32path, buf, 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
a, release, err := MapFloat32s(f32path, 0, n)
|
||||
if err != nil {
|
||||
t.Fatalf("MapFloat32s: %v", err)
|
||||
}
|
||||
if a.RawFloat32s()[42] != want32[42] {
|
||||
t.Fatalf("f32[42] = %v, want %v", a.RawFloat32s()[42], want32[42])
|
||||
}
|
||||
if err := release(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
ipath := filepath.Join(dir, "i64.bin")
|
||||
var ibuf []byte
|
||||
var iword [8]byte
|
||||
wantInt := make([]int64, n)
|
||||
for i := range n {
|
||||
wantInt[i] = int64(i) * 1_000_000
|
||||
binary.NativeEndian.PutUint64(iword[:], uint64(wantInt[i]))
|
||||
ibuf = append(ibuf, iword[:]...)
|
||||
}
|
||||
if err := os.WriteFile(ipath, ibuf, 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
b, releaseInt, err := MapInts(ipath, 0, n)
|
||||
if err != nil {
|
||||
t.Fatalf("MapInts: %v", err)
|
||||
}
|
||||
if b.RawInts()[13] != wantInt[13] {
|
||||
t.Fatalf("i64[13] = %d, want %d", b.RawInts()[13], wantInt[13])
|
||||
}
|
||||
if err := releaseInt(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
|
||||
// TestMapFloatsOffset maps from a non-zero offset inside a file that
|
||||
// carries a small header first. The header is a multiple of the
|
||||
// element size: the typed view the reader builds is built with unsafe
|
||||
// and must be aligned.
|
||||
func TestMapFloatsOffset(t *testing.T) {
|
||||
path := filepath.Join(t.TempDir(), "headered.bin")
|
||||
header := []byte("TENSOR HEADER 24 BYTES!!")
|
||||
if len(header)%8 != 0 {
|
||||
t.Fatalf("the test header is %d bytes, want a multiple of 8", len(header))
|
||||
}
|
||||
values := []float64{3.5, 1.25, -9}
|
||||
var buf []byte
|
||||
buf = append(buf, header...)
|
||||
var word [8]byte
|
||||
for _, v := range values {
|
||||
binary.NativeEndian.PutUint64(word[:], math.Float64bits(v))
|
||||
buf = append(buf, word[:]...)
|
||||
}
|
||||
if err := os.WriteFile(path, buf, 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
a, release, err := MapFloats(path, int64(len(header)), 3)
|
||||
if err != nil {
|
||||
t.Fatalf("MapFloats: %v", err)
|
||||
}
|
||||
for i, want := range values {
|
||||
if a.FloatAt(i) != want {
|
||||
t.Fatalf("value %d = %.14g, want %.14g", i, a.FloatAt(i), want)
|
||||
}
|
||||
}
|
||||
if err := release(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
|
||||
// TestMapErrors pins the validation contract: short files, negative
|
||||
// offsets, zero counts and directories are errors, never mappings.
|
||||
func TestMapErrors(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
path := filepath.Join(dir, "small.bin")
|
||||
if err := SaveNativeFloats(path, []float64{1, 2, 3}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if _, _, err := MapFloats(path, 4, 1); err == nil {
|
||||
t.Fatal("expected an error for an offset that misaligns the elements")
|
||||
}
|
||||
if _, _, err := MapFloats(path, 0, 4); err == nil {
|
||||
t.Fatal("expected an error when the file is too short")
|
||||
}
|
||||
if _, _, err := MapFloats(path, -8, 2); err == nil {
|
||||
t.Fatal("expected an error for a negative offset")
|
||||
}
|
||||
if _, _, err := MapFloats(path, 0, 0); err == nil {
|
||||
t.Fatal("expected an error for a zero count")
|
||||
}
|
||||
if _, _, err := MapFloats(dir, 0, 1); err == nil {
|
||||
t.Fatal("expected an error for a directory")
|
||||
}
|
||||
if _, _, err := MapFloats(filepath.Join(dir, "missing.bin"), 0, 1); err == nil {
|
||||
t.Fatal("expected an error for a missing file")
|
||||
}
|
||||
}
|
||||
+1077
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,102 @@
|
||||
// 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)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,634 @@
|
||||
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
|
||||
// SPDX-License-Identifier: MIT
|
||||
|
||||
package io
|
||||
|
||||
import (
|
||||
"encoding/hex"
|
||||
"math"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"slices"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"sourcedock.dev/petrbalvin/tensor/internal/core"
|
||||
)
|
||||
|
||||
func ncTempPath(t *testing.T) string {
|
||||
t.Helper()
|
||||
return filepath.Join(t.TempDir(), "data.nc")
|
||||
}
|
||||
|
||||
// TestNetCDFRoundTrip moves two variables of different dtypes through
|
||||
// the file format and back, values, shapes, dimensions and attributes
|
||||
// included.
|
||||
func TestNetCDFRoundTrip(t *testing.T) {
|
||||
dims := []NetCDFDim{{Name: "lat", Length: 3}, {Name: "lon", Length: 4}}
|
||||
temp := mustFloats(t, []float64{
|
||||
1.5, -2.25, 3.125, 4,
|
||||
-5.5, 6.75, -7.875, 8,
|
||||
9.25, -10.5, 11.125, -12,
|
||||
}, 3, 4)
|
||||
// A scalar carries one value and no dimensions.
|
||||
scalar, err := core.FromFloats([]float64{42}, 1)
|
||||
if err != nil {
|
||||
t.Fatalf("FromFloats: %v", err)
|
||||
}
|
||||
vars := []NetCDFVar{
|
||||
{
|
||||
Name: "temp",
|
||||
Dims: []string{"lat", "lon"},
|
||||
Values: temp,
|
||||
Attrs: map[string]string{"units": "degC", "long_name": "sea surface"},
|
||||
},
|
||||
{
|
||||
Name: "quality",
|
||||
Values: scalar,
|
||||
},
|
||||
}
|
||||
|
||||
path := ncTempPath(t)
|
||||
attrs := map[string]string{"title": "round trip", "source": "tensor"}
|
||||
if err := SaveNetCDF(path, dims, vars, attrs); err != nil {
|
||||
t.Fatalf("SaveNetCDF: %v", err)
|
||||
}
|
||||
gotDims, gotVars, gotAttrs, err := LoadNetCDF(path)
|
||||
if err != nil {
|
||||
t.Fatalf("LoadNetCDF: %v", err)
|
||||
}
|
||||
if len(gotDims) != 2 || gotDims[0] != dims[0] || gotDims[1] != dims[1] {
|
||||
t.Fatalf("dims = %v, want %v", gotDims, dims)
|
||||
}
|
||||
for k, want := range attrs {
|
||||
if gotAttrs[k] != want {
|
||||
t.Fatalf("global attr %q = %q, want %q", k, gotAttrs[k], want)
|
||||
}
|
||||
}
|
||||
if len(gotVars) != 2 {
|
||||
t.Fatalf("got %d variables, want 2", len(gotVars))
|
||||
}
|
||||
tv := gotVars[0]
|
||||
if tv.Name != "temp" || len(tv.Dims) != 2 || tv.Dims[0] != "lat" || tv.Dims[1] != "lon" {
|
||||
t.Fatalf("temp dims = %v", tv.Dims)
|
||||
}
|
||||
if tv.Values.NDim() != 2 || tv.Values.Shape()[0] != 3 || tv.Values.Shape()[1] != 4 {
|
||||
t.Fatalf("temp shape = %v", tv.Values.Shape())
|
||||
}
|
||||
for i := range temp.Len() {
|
||||
if tv.Values.FloatAt(i) != temp.FloatAt(i) {
|
||||
t.Fatalf("temp[%d] = %g, want %g", i, tv.Values.FloatAt(i), temp.FloatAt(i))
|
||||
}
|
||||
}
|
||||
for k, want := range map[string]string{"units": "degC", "long_name": "sea surface"} {
|
||||
if tv.Attrs[k] != want {
|
||||
t.Fatalf("temp attr %q = %q, want %q", k, tv.Attrs[k], want)
|
||||
}
|
||||
}
|
||||
qv := gotVars[1]
|
||||
if qv.Name != "quality" || len(qv.Dims) != 0 || qv.Values.FloatAt(0) != 42 {
|
||||
t.Fatalf("quality = %v %v %g", qv.Name, qv.Dims, qv.Values.FloatAt(0))
|
||||
}
|
||||
}
|
||||
|
||||
// TestNetCDFIntAndFloat32Ranges pins the integer and single-precision
|
||||
// external types, extremes included: NC_INT is a signed 32-bit value
|
||||
// that lands int32 on read, and the float32 round trip is exact
|
||||
// through the float64 landing NC_FLOAT keeps.
|
||||
func TestNetCDFIntAndFloat32Ranges(t *testing.T) {
|
||||
ints, err := core.FromInts([]int64{-2147483648, -1, 0, 1, 2147483647}, 5)
|
||||
if err != nil {
|
||||
t.Fatalf("FromInts: %v", err)
|
||||
}
|
||||
f32, err := core.FromFloat32s([]float32{1.5, -2.25, 0, math.MaxFloat32, math.SmallestNonzeroFloat32}, 5)
|
||||
if err != nil {
|
||||
t.Fatalf("FromFloat32s: %v", err)
|
||||
}
|
||||
path := ncTempPath(t)
|
||||
vars := []NetCDFVar{
|
||||
{Name: "i", Dims: []string{"n"}, Values: ints},
|
||||
{Name: "f", Dims: []string{"n"}, Values: f32},
|
||||
}
|
||||
if err := SaveNetCDF(path, []NetCDFDim{{Name: "n", Length: 5}}, vars, nil); err != nil {
|
||||
t.Fatalf("SaveNetCDF: %v", err)
|
||||
}
|
||||
_, got, _, err := LoadNetCDF(path)
|
||||
if err != nil {
|
||||
t.Fatalf("LoadNetCDF: %v", err)
|
||||
}
|
||||
if dt := got[0].Values.Dtype(); dt != core.Int32 {
|
||||
t.Fatalf("int dtype = %s, want int32", dt)
|
||||
}
|
||||
wantInts := []int32{-2147483648, -1, 0, 1, 2147483647}
|
||||
for i, want := range wantInts {
|
||||
if got[0].Values.RawInt32s()[i] != want {
|
||||
t.Fatalf("int[%d] = %d, want %d", i, got[0].Values.RawInt32s()[i], want)
|
||||
}
|
||||
}
|
||||
if dt := got[1].Values.Dtype(); dt != core.Float {
|
||||
t.Fatalf("float32 dtype = %s, want the float64 landing NC_FLOAT keeps", dt)
|
||||
}
|
||||
for i := range 5 {
|
||||
if got[1].Values.FloatAt(i) != float64(f32.RawFloat32s()[i]) {
|
||||
t.Fatalf("f32[%d] = %g, want %g", i, got[1].Values.FloatAt(i), f32.RawFloat32s()[i])
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// TestNetCDFParsesKnownBytes pins the parser against the specification
|
||||
// rather than against our own writer: this is the canonical tiny.nc
|
||||
// from the NetCDF file format documentation, one SHORT variable "vx"
|
||||
// of length 5 with values 3, 1, 4, 1, 5 and one fill-padded tail.
|
||||
func TestNetCDFParsesKnownBytes(t *testing.T) {
|
||||
const dump = "" +
|
||||
"43444601" + // magic CDF\x01
|
||||
"00000000" + // numrecs = 0
|
||||
"0000000a" + // NC_DIMENSION
|
||||
"00000001" + // one dimension
|
||||
"0000000364696d00" + // name "dim" + pad
|
||||
"00000005" + // length 5
|
||||
"0000000000000000" + // gatt_list ABSENT
|
||||
"0000000b" + // NC_VARIABLE
|
||||
"00000001" + // one variable
|
||||
"0000000276780000" + // name "vx" + pad
|
||||
"00000001" + // rank 1
|
||||
"00000000" + // dimid 0
|
||||
"0000000000000000" + // vatt_list ABSENT
|
||||
"00000003" + // NC_SHORT
|
||||
"0000000c" + // vsize 12 (5 shorts + 1 fill pad)
|
||||
"00000050" + // begin 80
|
||||
"00030001000400010005" + // 3, 1, 4, 1, 5
|
||||
"8001" // fill pad
|
||||
raw, err := hex.DecodeString(dump)
|
||||
if err != nil {
|
||||
t.Fatalf("hex: %v", err)
|
||||
}
|
||||
if len(raw) != 92 {
|
||||
t.Fatalf("fixture is %d bytes, the specification says 92", len(raw))
|
||||
}
|
||||
path := ncTempPath(t)
|
||||
if err := os.WriteFile(path, raw, 0o644); err != nil {
|
||||
t.Fatalf("write: %v", err)
|
||||
}
|
||||
dims, vars, attrs, err := LoadNetCDF(path)
|
||||
if err != nil {
|
||||
t.Fatalf("LoadNetCDF: %v", err)
|
||||
}
|
||||
if len(dims) != 1 || dims[0].Name != "dim" || dims[0].Length != 5 {
|
||||
t.Fatalf("dims = %v", dims)
|
||||
}
|
||||
if len(vars) != 1 || vars[0].Name != "vx" || len(vars[0].Dims) != 1 || vars[0].Dims[0] != "dim" {
|
||||
t.Fatalf("vars = %v", vars)
|
||||
}
|
||||
// NC_SHORT lands int16, the width the classic file stores.
|
||||
if dt := vars[0].Values.Dtype(); dt != core.Int16 {
|
||||
t.Fatalf("vx dtype = %s, want int16", dt)
|
||||
}
|
||||
if got := vars[0].Values.RawInt16s()[:5]; !slices.Equal(got, []int16{3, 1, 4, 1, 5}) {
|
||||
t.Fatalf("vx = %v, want [3 1 4 1 5]", got)
|
||||
}
|
||||
want := []float64{3, 1, 4, 1, 5}
|
||||
for i, w := range want {
|
||||
if got := vars[0].Values.FloatAt(i); got != w {
|
||||
t.Fatalf("vx[%d] = %g, want %g", i, got, w)
|
||||
}
|
||||
}
|
||||
if len(attrs) != 0 {
|
||||
t.Fatalf("attrs = %v, want none", attrs)
|
||||
}
|
||||
}
|
||||
|
||||
// TestNetCDFErrors pins the refusal paths: a record dimension, an
|
||||
// unknown version, a truncated file, bad names and a shape mismatch.
|
||||
func TestNetCDFErrors(t *testing.T) {
|
||||
path := ncTempPath(t)
|
||||
a := mustFloats(t, []float64{1, 2, 3}, 3)
|
||||
|
||||
// The record dimension is supported, but only where the classic
|
||||
// model allows it: leading the dimension list, leading each record
|
||||
// variable's dimensions, and holding a whole number of records.
|
||||
err := SaveNetCDF(path, []NetCDFDim{{Name: "x", Length: 2}, {Name: "t", Length: 0}},
|
||||
[]NetCDFVar{{Name: "v", Dims: []string{"x"}, Values: a}}, nil)
|
||||
if err == nil || !strings.Contains(err.Error(), "leads the list") {
|
||||
t.Fatalf("record dimension not first: %v", err)
|
||||
}
|
||||
err = SaveNetCDF(path, []NetCDFDim{{Name: "t", Length: 0}, {Name: "x", Length: 3}},
|
||||
[]NetCDFVar{{Name: "v", Dims: []string{"x", "t"}, Values: a}}, nil)
|
||||
if err == nil || !strings.Contains(err.Error(), "record axis leads") {
|
||||
t.Fatalf("record axis not leading: %v", err)
|
||||
}
|
||||
err = SaveNetCDF(path, []NetCDFDim{{Name: "t", Length: 0}, {Name: "x", Length: 2}},
|
||||
[]NetCDFVar{{Name: "v", Dims: []string{"t", "x"}, Values: a}}, nil)
|
||||
if err == nil || !strings.Contains(err.Error(), "whole number of records") {
|
||||
t.Fatalf("partial record: %v", err)
|
||||
}
|
||||
two, terr := core.FromFloats([]float64{1, 2}, 2)
|
||||
if terr != nil {
|
||||
t.Fatalf("FromFloats: %v", terr)
|
||||
}
|
||||
err = SaveNetCDF(path, []NetCDFDim{{Name: "t", Length: 0}, {Name: "x", Length: 1}},
|
||||
[]NetCDFVar{
|
||||
{Name: "u", Dims: []string{"t"}, Values: a},
|
||||
{Name: "w", Dims: []string{"t", "x"}, Values: two},
|
||||
}, nil)
|
||||
if err == nil || !strings.Contains(err.Error(), "records, the others") {
|
||||
t.Fatalf("record count disagreement: %v", err)
|
||||
}
|
||||
|
||||
// An int64 value outside int32 is refused, not truncated.
|
||||
big, ferr := core.FromInts([]int64{1 << 40}, 1)
|
||||
if ferr != nil {
|
||||
t.Fatalf("FromInts: %v", ferr)
|
||||
}
|
||||
err = SaveNetCDF(path, []NetCDFDim{{Name: "n", Length: 1}},
|
||||
[]NetCDFVar{{Name: "v", Dims: []string{"n"}, Values: big}}, nil)
|
||||
if err == nil || !strings.Contains(err.Error(), "int32 range") {
|
||||
t.Fatalf("int64 out of range: %v", err)
|
||||
}
|
||||
|
||||
// A landed narrow variable is refused by the writer, whose classic
|
||||
// model stores float64, float32 and int64 arrays only.
|
||||
i8, ierr := core.FromInt8s([]int8{1, 2, 3}, 3)
|
||||
if ierr != nil {
|
||||
t.Fatalf("FromInt8s: %v", ierr)
|
||||
}
|
||||
err = SaveNetCDF(path, []NetCDFDim{{Name: "n", Length: 3}},
|
||||
[]NetCDFVar{{Name: "v", Dims: []string{"n"}, Values: i8}}, nil)
|
||||
if err == nil || !strings.Contains(err.Error(), "SaveNetCDF") ||
|
||||
!strings.Contains(err.Error(), "int8") ||
|
||||
!strings.Contains(err.Error(), "the classic model stores float64, float32 and int arrays") {
|
||||
t.Fatalf("int8 variable: %v", err)
|
||||
}
|
||||
|
||||
// The value count must match the dimensions.
|
||||
err = SaveNetCDF(path, []NetCDFDim{{Name: "n", Length: 4}},
|
||||
[]NetCDFVar{{Name: "v", Dims: []string{"n"}, Values: a}}, nil)
|
||||
if err == nil || !strings.Contains(err.Error(), "its dimensions hold 4") {
|
||||
t.Fatalf("shape mismatch: %v", err)
|
||||
}
|
||||
|
||||
// Names follow the traditional grammar.
|
||||
err = SaveNetCDF(path, []NetCDFDim{{Name: "bad name", Length: 1}},
|
||||
[]NetCDFVar{{Name: "v", Dims: []string{"bad name"}, Values: mustFloats(t, []float64{1}, 1)}}, nil)
|
||||
if err == nil || !strings.Contains(err.Error(), "forbids") {
|
||||
t.Fatalf("bad name: %v", err)
|
||||
}
|
||||
|
||||
// Duplicate variable names are refused.
|
||||
err = SaveNetCDF(path, []NetCDFDim{{Name: "n", Length: 1}},
|
||||
[]NetCDFVar{
|
||||
{Name: "v", Dims: []string{"n"}, Values: mustFloats(t, []float64{1}, 1)},
|
||||
{Name: "v", Dims: []string{"n"}, Values: mustFloats(t, []float64{2}, 1)},
|
||||
}, nil)
|
||||
if err == nil || !strings.Contains(err.Error(), "declared twice") {
|
||||
t.Fatalf("duplicate variable: %v", err)
|
||||
}
|
||||
|
||||
// Not a NetCDF file at all.
|
||||
if err := os.WriteFile(path, []byte("not a netcdf file at all, sorry"), 0o644); err != nil {
|
||||
t.Fatalf("write: %v", err)
|
||||
}
|
||||
if _, _, _, err := LoadNetCDF(path); err == nil || !strings.Contains(err.Error(), "not a NetCDF file") {
|
||||
t.Fatalf("bad magic: %v", err)
|
||||
}
|
||||
|
||||
// A version this module does not speak.
|
||||
if err := os.WriteFile(path, []byte{'C', 'D', 'F', 9}, 0o644); err != nil {
|
||||
t.Fatalf("write: %v", err)
|
||||
}
|
||||
if _, _, _, err := LoadNetCDF(path); err == nil || !strings.Contains(err.Error(), "unsupported NetCDF version") {
|
||||
t.Fatalf("bad version: %v", err)
|
||||
}
|
||||
|
||||
// CDF-5, the 64-bit-count extension: refused by the same version
|
||||
// gate rather than parsed as classic.
|
||||
if err := os.WriteFile(path, []byte{'C', 'D', 'F', 5}, 0o644); err != nil {
|
||||
t.Fatalf("write: %v", err)
|
||||
}
|
||||
if _, _, _, err := LoadNetCDF(path); err == nil || !strings.Contains(err.Error(), "unsupported NetCDF version") {
|
||||
t.Fatalf("CDF-5 version: %v", err)
|
||||
}
|
||||
|
||||
// A truncated header.
|
||||
if err := os.WriteFile(path, []byte{'C', 'D', 'F', 1, 0, 0}, 0o644); err != nil {
|
||||
t.Fatalf("write: %v", err)
|
||||
}
|
||||
if _, _, _, err := LoadNetCDF(path); err == nil || !strings.Contains(err.Error(), "ends inside the header") {
|
||||
t.Fatalf("truncated: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
// TestNetCDFRecordDimensionRead pins the record dimension at load
|
||||
// time: a length of 0 in the leading dimension is the unlimited
|
||||
// dimension, so the header parses and the dimension comes back with
|
||||
// length 0 (its extent lives on each record variable's first axis).
|
||||
func TestNetCDFRecordDimensionRead(t *testing.T) {
|
||||
const dump = "" +
|
||||
"43444601" + // magic
|
||||
"00000001" + // numrecs = 1
|
||||
"0000000a00000001" + // NC_DIMENSION, one dimension
|
||||
"0000000174000000" + // name "t" + pad
|
||||
"00000000" + // length 0: the record dimension
|
||||
"0000000000000000" + // gatt ABSENT
|
||||
"0000000b00000000" // NC_VARIABLE, zero variables
|
||||
raw, err := hex.DecodeString(dump)
|
||||
if err != nil {
|
||||
t.Fatalf("hex: %v", err)
|
||||
}
|
||||
path := ncTempPath(t)
|
||||
if err := os.WriteFile(path, raw, 0o644); err != nil {
|
||||
t.Fatalf("write: %v", err)
|
||||
}
|
||||
dims, vars, _, err := LoadNetCDF(path)
|
||||
if err != nil {
|
||||
t.Fatalf("record dimension: %v", err)
|
||||
}
|
||||
if len(dims) != 1 || dims[0].Name != "t" || dims[0].Length != 0 {
|
||||
t.Fatalf("dims = %+v, want one record dimension named t with length 0", dims)
|
||||
}
|
||||
if len(vars) != 0 {
|
||||
t.Fatalf("vars = %+v, want none", vars)
|
||||
}
|
||||
// A second zero-length dimension is refused: only one unlimited
|
||||
// dimension exists, and it leads the list.
|
||||
const twoRecords = "" +
|
||||
"43444601" + // magic
|
||||
"00000001" + // numrecs = 1
|
||||
"0000000a00000002" + // NC_DIMENSION, two dimensions
|
||||
"0000000174000000" + // name "t" + pad
|
||||
"00000000" + // length 0: the record dimension
|
||||
"0000000178000000" + // name "x" + pad
|
||||
"00000000" + // length 0 again: refused
|
||||
"0000000000000000" + // gatt ABSENT
|
||||
"0000000b00000000" // NC_VARIABLE, zero variables
|
||||
raw, err = hex.DecodeString(twoRecords)
|
||||
if err != nil {
|
||||
t.Fatalf("hex: %v", err)
|
||||
}
|
||||
if err := os.WriteFile(path, raw, 0o644); err != nil {
|
||||
t.Fatalf("write: %v", err)
|
||||
}
|
||||
if _, _, _, err := LoadNetCDF(path); err == nil {
|
||||
t.Fatal("expected an error for a second zero-length dimension")
|
||||
}
|
||||
}
|
||||
|
||||
// TestNetCDFClassicNativeLandings pins the classic type codes onto the
|
||||
// core dtypes they land, against a hand-built CDF-1 file rather than
|
||||
// the package's own writer: NC_BYTE as int8, NC_SHORT as int16,
|
||||
// NC_INT as int32, NC_CHAR as uint8 raw bytes, and NC_FLOAT and
|
||||
// NC_DOUBLE keeping their float64 landing. The values carry the
|
||||
// extremes of every width, so a widened or unsigned reading of any of
|
||||
// them fails the pin.
|
||||
func TestNetCDFClassicNativeLandings(t *testing.T) {
|
||||
const dump = "" +
|
||||
"43444601" + // magic CDF\x01
|
||||
"00000000" + // numrecs = 0
|
||||
"0000000a00000001" + // NC_DIMENSION, one dimension
|
||||
"000000016e000000" + // name "n" + pad
|
||||
"00000004" + // length 4
|
||||
"0000000000000000" + // gatt_list ABSENT
|
||||
"0000000b00000006" + // NC_VARIABLE, six variables
|
||||
"000000016200000000000001000000000000000000000000" + // "b": rank 1, dim 0, no attrs
|
||||
"00000001" + // NC_BYTE
|
||||
"00000004" + // vsize 4
|
||||
"00000104" + // begin 260
|
||||
"000000017300000000000001000000000000000000000000" + // "s"
|
||||
"00000003" + // NC_SHORT
|
||||
"00000008" +
|
||||
"00000108" + // begin 264
|
||||
"000000016900000000000001000000000000000000000000" + // "i"
|
||||
"00000004" + // NC_INT
|
||||
"00000010" +
|
||||
"00000110" + // begin 272
|
||||
"000000016300000000000001000000000000000000000000" + // "c"
|
||||
"00000002" + // NC_CHAR
|
||||
"00000004" +
|
||||
"00000120" + // begin 288
|
||||
"000000016600000000000001000000000000000000000000" + // "f"
|
||||
"00000005" + // NC_FLOAT
|
||||
"00000010" +
|
||||
"00000124" + // begin 292
|
||||
"000000016400000000000001000000000000000000000000" + // "d"
|
||||
"00000006" + // NC_DOUBLE
|
||||
"00000020" +
|
||||
"00000134" + // begin 308
|
||||
"80ff007f" + // b: -128, -1, 0, 127
|
||||
"8000ffff00017fff" + // s: -32768, -1, 1, 32767
|
||||
"80000000ffffffff000000017fffffff" + // i: int32 extremes
|
||||
"0041c8ff" + // c: raw bytes 0, 65, 200, 255
|
||||
"3fc00000c02000000000000040500000" + // f: 1.5, -2.5, 0, 3.25
|
||||
"3ff8000000000000c004000000000000" + // d: 1.5, -2.5,
|
||||
"00000000000000004010000000000000" // 0, 4
|
||||
raw, err := hex.DecodeString(dump)
|
||||
if err != nil {
|
||||
t.Fatalf("hex: %v", err)
|
||||
}
|
||||
path := ncTempPath(t)
|
||||
if err := os.WriteFile(path, raw, 0o644); err != nil {
|
||||
t.Fatalf("write: %v", err)
|
||||
}
|
||||
dims, vars, attrs, err := LoadNetCDF(path)
|
||||
if err != nil {
|
||||
t.Fatalf("LoadNetCDF: %v", err)
|
||||
}
|
||||
if len(dims) != 1 || dims[0].Name != "n" || dims[0].Length != 4 {
|
||||
t.Fatalf("dims = %v", dims)
|
||||
}
|
||||
if len(attrs) != 0 {
|
||||
t.Fatalf("attrs = %v, want none", attrs)
|
||||
}
|
||||
if len(vars) != 6 {
|
||||
t.Fatalf("vars = %d, want 6", len(vars))
|
||||
}
|
||||
byName := map[string]NetCDFVar{}
|
||||
for _, v := range vars {
|
||||
byName[v.Name] = v
|
||||
if s := v.Values.Shape(); len(s) != 1 || s[0] != 4 {
|
||||
t.Fatalf("%s shape = %v, want [4]", v.Name, s)
|
||||
}
|
||||
}
|
||||
if dt := byName["b"].Values.Dtype(); dt != core.Int8 {
|
||||
t.Fatalf("NC_BYTE dtype = %s, want int8", dt)
|
||||
}
|
||||
if got, want := byName["b"].Values.RawInt8s()[:4], []int8{-128, -1, 0, 127}; !slices.Equal(got, want) {
|
||||
t.Fatalf("NC_BYTE values = %v, want %v", got, want)
|
||||
}
|
||||
if dt := byName["s"].Values.Dtype(); dt != core.Int16 {
|
||||
t.Fatalf("NC_SHORT dtype = %s, want int16", dt)
|
||||
}
|
||||
if got, want := byName["s"].Values.RawInt16s()[:4], []int16{-32768, -1, 1, 32767}; !slices.Equal(got, want) {
|
||||
t.Fatalf("NC_SHORT values = %v, want %v", got, want)
|
||||
}
|
||||
if dt := byName["i"].Values.Dtype(); dt != core.Int32 {
|
||||
t.Fatalf("NC_INT dtype = %s, want int32", dt)
|
||||
}
|
||||
if got, want := byName["i"].Values.RawInt32s()[:4], []int32{-2147483648, -1, 1, 2147483647}; !slices.Equal(got, want) {
|
||||
t.Fatalf("NC_INT values = %v, want %v", got, want)
|
||||
}
|
||||
// CHAR carries bytes, not text: the array level lands uint8 and
|
||||
// what the bytes spell is the caller's question.
|
||||
if dt := byName["c"].Values.Dtype(); dt != core.Uint8 {
|
||||
t.Fatalf("NC_CHAR dtype = %s, want uint8", dt)
|
||||
}
|
||||
if got, want := byName["c"].Values.RawUint8s()[:4], []uint8{0, 65, 200, 255}; !slices.Equal(got, want) {
|
||||
t.Fatalf("NC_CHAR values = %v, want %v", got, want)
|
||||
}
|
||||
if dt := byName["f"].Values.Dtype(); dt != core.Float {
|
||||
t.Fatalf("NC_FLOAT dtype = %s, want float", dt)
|
||||
}
|
||||
for i, want := range []float64{1.5, -2.5, 0, 3.25} {
|
||||
if got := byName["f"].Values.FloatAt(i); got != want {
|
||||
t.Fatalf("f[%d] = %v, want %v", i, got, want)
|
||||
}
|
||||
}
|
||||
if dt := byName["d"].Values.Dtype(); dt != core.Float {
|
||||
t.Fatalf("NC_DOUBLE dtype = %s, want float", dt)
|
||||
}
|
||||
for i, want := range []float64{1.5, -2.5, 0, 4} {
|
||||
if got := byName["d"].Values.FloatAt(i); got != want {
|
||||
t.Fatalf("d[%d] = %v, want %v", i, got, want)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// TestNetCDFRecordNarrowLandings pins the record path on a hand-built
|
||||
// CDF-1 file: an NC_BYTE record variable's slabs land int8 in place
|
||||
// and an NC_SHORT record variable's slabs land int16, including the
|
||||
// slab padding the classic layout writes between records, and a fixed
|
||||
// NC_SHORT variable lands int16 beside them.
|
||||
func TestNetCDFRecordNarrowLandings(t *testing.T) {
|
||||
const dump = "" +
|
||||
"43444601" + // magic
|
||||
"00000002" + // numrecs = 2
|
||||
"0000000a00000002" + // NC_DIMENSION, two dimensions
|
||||
"0000000372656300" + // name "rec" + pad
|
||||
"00000000" + // length 0: the record dimension
|
||||
"0000000178000000" + // name "x" + pad
|
||||
"00000003" + // length 3
|
||||
"0000000000000000" + // gatt_list ABSENT
|
||||
"0000000b00000003" + // NC_VARIABLE, three variables
|
||||
"0000000272620000" + // name "rb" + pad
|
||||
"00000001" + // rank 1
|
||||
"00000000" + // dimid 0: the record dimension
|
||||
"0000000000000000" + // vatt ABSENT
|
||||
"00000001" + // NC_BYTE
|
||||
"00000004" + // vsize 4: one byte padded to four
|
||||
"000000ac" + // begin 172: its slab inside the first record
|
||||
"0000000273730000" + // name "ss" + pad
|
||||
"00000001" + // rank 1
|
||||
"00000001" + // dimid 1: x, fixed
|
||||
"0000000000000000" + // vatt ABSENT
|
||||
"00000003" + // NC_SHORT
|
||||
"00000008" + // vsize 8: six bytes padded to eight
|
||||
"000000a4" + // begin 164
|
||||
"0000000272730000" + // name "rs" + pad
|
||||
"00000001" + // rank 1
|
||||
"00000000" + // dimid 0: the record dimension
|
||||
"0000000000000000" + // vatt ABSENT
|
||||
"00000003" + // NC_SHORT
|
||||
"00000004" + // vsize 4: one short padded to four
|
||||
"000000b0" + // begin 176: its slab inside the first record
|
||||
"0005fffa0007" + // ss: 5, -6, 7
|
||||
"0000" + // slab padding
|
||||
"07000000" + // rb record 0: 7
|
||||
"12340000" + // rs record 0: 0x1234
|
||||
"c8000000" + // rb record 1: 200 as int8 is -56
|
||||
"fff00000" // rs record 1: -16
|
||||
raw, err := hex.DecodeString(dump)
|
||||
if err != nil {
|
||||
t.Fatalf("hex: %v", err)
|
||||
}
|
||||
path := ncTempPath(t)
|
||||
if err := os.WriteFile(path, raw, 0o644); err != nil {
|
||||
t.Fatalf("write: %v", err)
|
||||
}
|
||||
dims, vars, _, err := LoadNetCDF(path)
|
||||
if err != nil {
|
||||
t.Fatalf("LoadNetCDF: %v", err)
|
||||
}
|
||||
if len(dims) != 2 || dims[0].Name != "rec" || dims[0].Length != 0 || dims[1].Length != 3 {
|
||||
t.Fatalf("dims = %v", dims)
|
||||
}
|
||||
byName := map[string]NetCDFVar{}
|
||||
for _, v := range vars {
|
||||
byName[v.Name] = v
|
||||
}
|
||||
rb := byName["rb"].Values
|
||||
if rb.Dtype() != core.Int8 {
|
||||
t.Fatalf("rb dtype = %s, want int8", rb.Dtype())
|
||||
}
|
||||
if s := rb.Shape(); len(s) != 1 || s[0] != 2 {
|
||||
t.Fatalf("rb shape = %v, want [2]: the record axis carries the record count", s)
|
||||
}
|
||||
if got, want := rb.RawInt8s()[:2], []int8{7, -56}; !slices.Equal(got, want) {
|
||||
t.Fatalf("rb values = %v, want %v", got, want)
|
||||
}
|
||||
rs := byName["rs"].Values
|
||||
if rs.Dtype() != core.Int16 {
|
||||
t.Fatalf("rs dtype = %s, want int16", rs.Dtype())
|
||||
}
|
||||
if s := rs.Shape(); len(s) != 1 || s[0] != 2 {
|
||||
t.Fatalf("rs shape = %v, want [2]: the record axis carries the record count", s)
|
||||
}
|
||||
if got, want := rs.RawInt16s()[:2], []int16{0x1234, -16}; !slices.Equal(got, want) {
|
||||
t.Fatalf("rs values = %v, want %v", got, want)
|
||||
}
|
||||
ss := byName["ss"].Values
|
||||
if ss.Dtype() != core.Int16 {
|
||||
t.Fatalf("ss dtype = %s, want int16", ss.Dtype())
|
||||
}
|
||||
if got, want := ss.RawInt16s()[:3], []int16{5, -6, 7}; !slices.Equal(got, want) {
|
||||
t.Fatalf("ss values = %v, want %v", got, want)
|
||||
}
|
||||
}
|
||||
|
||||
// TestNetCDFUnknownTypeRefused pins the refusal of every type code
|
||||
// beyond the classic six: the unsigned and 64-bit codes exist only in
|
||||
// formats this module does not speak, and a header carrying one is
|
||||
// refused by name rather than decoded into some other dtype.
|
||||
func TestNetCDFUnknownTypeRefused(t *testing.T) {
|
||||
const template = "" +
|
||||
"43444601" + // magic
|
||||
"00000000" + // numrecs = 0
|
||||
"0000000000000000" + // dim_list ABSENT
|
||||
"0000000000000000" + // gatt_list ABSENT
|
||||
"0000000b" + // NC_VARIABLE
|
||||
"00000001" + // one variable
|
||||
"0000000176000000" + // name "v" + pad
|
||||
"00000000" + // rank 0
|
||||
"0000000000000000" + // vatt ABSENT
|
||||
"00000007" + // type code, replaced per case
|
||||
"00000000" + // vsize
|
||||
"00000000" // begin
|
||||
for _, tc := range []struct{ code, word string }{
|
||||
{"7", "00000007"}, {"8", "00000008"}, {"9", "00000009"},
|
||||
{"10", "0000000a"}, {"11", "0000000b"},
|
||||
} {
|
||||
raw, err := hex.DecodeString(strings.Replace(template, "00000007", tc.word, 1))
|
||||
if err != nil {
|
||||
t.Fatalf("hex: %v", err)
|
||||
}
|
||||
path := ncTempPath(t)
|
||||
if err := os.WriteFile(path, raw, 0o644); err != nil {
|
||||
t.Fatalf("write: %v", err)
|
||||
}
|
||||
_, _, _, err = LoadNetCDF(path)
|
||||
if err == nil || !strings.Contains(err.Error(), "unknown type "+tc.code) {
|
||||
t.Fatalf("type code %s: err = %v, want the unknown-type refusal", tc.code, err)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// TestNetCDFDecodeCellsRefused pins the explicit refusal inside
|
||||
// decodeCells itself: an ncType without a case is a named error, never
|
||||
// a buffer left silently untouched.
|
||||
func TestNetCDFDecodeCellsRefused(t *testing.T) {
|
||||
arr := core.New(core.Float, 2)
|
||||
for _, code := range []uint32{0, 7, 8, 9, 10, 11, 12, 99} {
|
||||
err := decodeCells(arr, make([]byte, 16), code, 0, 2)
|
||||
if err == nil || !strings.Contains(err.Error(), "unknown type") {
|
||||
t.Fatalf("decodeCells with type %d: err = %v, want the unknown-type refusal", code, err)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,158 @@
|
||||
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
|
||||
// SPDX-License-Identifier: MIT
|
||||
|
||||
package io
|
||||
|
||||
import (
|
||||
"encoding/hex"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"testing"
|
||||
|
||||
"sourcedock.dev/petrbalvin/tensor/internal/core"
|
||||
)
|
||||
|
||||
// The record (unlimited) dimension: written and read back, and read
|
||||
// from a file another implementation wrote.
|
||||
|
||||
// TestNetCDFRecordRoundTrip writes a file with a record dimension and a
|
||||
// fixed variable, then reads it back: the record count, the shapes and
|
||||
// the values must survive, and the record axis must stay first.
|
||||
func TestNetCDFRecordRoundTrip(t *testing.T) {
|
||||
path := filepath.Join(t.TempDir(), "rec.nc")
|
||||
tp := mustFloats(t, []float64{1.5, 2.5, 3.5}, 3)
|
||||
sp := mustInts(t, []int64{1, 2, 3, 4, 5, 6, 7, 8, 9}, 3, 3)
|
||||
vp := mustFloats32(t, []float32{7, 8, 9}, 3)
|
||||
dims := []NetCDFDim{{Name: "time", Length: 0}, {Name: "x", Length: 3}}
|
||||
vars := []NetCDFVar{
|
||||
{Name: "t", Dims: []string{"time"}, Values: tp},
|
||||
{Name: "s", Dims: []string{"time", "x"}, Values: sp},
|
||||
{Name: "v", Dims: []string{"x"}, Values: vp},
|
||||
}
|
||||
if err := SaveNetCDF(path, dims, vars, map[string]string{"title": "record round trip"}); err != nil {
|
||||
t.Fatalf("SaveNetCDF: %v", err)
|
||||
}
|
||||
gotDims, gotVars, attrs, err := LoadNetCDF(path)
|
||||
if err != nil {
|
||||
t.Fatalf("LoadNetCDF: %v", err)
|
||||
}
|
||||
if len(gotDims) != 2 || gotDims[0].Name != "time" || gotDims[0].Length != 0 || gotDims[1].Length != 3 {
|
||||
t.Fatalf("dims = %+v, want the record dimension first with length 0", gotDims)
|
||||
}
|
||||
if attrs["title"] != "record round trip" {
|
||||
t.Fatalf("attrs = %v", attrs)
|
||||
}
|
||||
if len(gotVars) != 3 {
|
||||
t.Fatalf("vars = %d, want 3", len(gotVars))
|
||||
}
|
||||
byName := map[string]NetCDFVar{}
|
||||
for _, v := range gotVars {
|
||||
byName[v.Name] = v
|
||||
}
|
||||
if s := byName["t"].Values.Shape(); s[0] != 3 {
|
||||
t.Fatalf("t shape = %v, want [3]", s)
|
||||
}
|
||||
if got := byName["s"].Values.Shape(); got[0] != 3 || got[1] != 3 {
|
||||
t.Fatalf("s shape = %v, want [3 3]", got)
|
||||
}
|
||||
for i, want := range []float64{1.5, 2.5, 3.5} {
|
||||
if got := byName["t"].Values.FloatAt(i); got != want {
|
||||
t.Fatalf("t[%d] = %v, want %v", i, got, want)
|
||||
}
|
||||
}
|
||||
for i, want := range []float64{1, 2, 3, 4, 5, 6, 7, 8, 9} {
|
||||
if got := byName["s"].Values.FloatAt(i); got != want {
|
||||
t.Fatalf("s[%d] = %v, want %v", i, got, want)
|
||||
}
|
||||
}
|
||||
for i, want := range []float64{7, 8, 9} {
|
||||
if got := byName["v"].Values.FloatAt(i); got != want {
|
||||
t.Fatalf("v[%d] = %v, want %v", i, got, want)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// TestNetCDFRecordFromNetCDF4 reads a file the netCDF reference
|
||||
// implementation wrote: an unlimited time axis, a 1-D record variable
|
||||
// of NC_DOUBLE, a 2-D record variable of NC_SHORT whose six-byte slab
|
||||
// is padded to eight on disk, and a fixed NC_FLOAT variable. It pins
|
||||
// the layout a foreign writer produces, padding included.
|
||||
func TestNetCDFRecordFromNetCDF4(t *testing.T) {
|
||||
const dump = "" +
|
||||
"43444601000000020000000a000000020000000474696d6500000000000000017800000000000003" +
|
||||
"00000000000000000000000b00000003000000017400000000000001000000000000000000000000" +
|
||||
"0000000600000008000000b400000001730000000000000200000000000000010000000000000000" +
|
||||
"0000000300000008000000bc00000001760000000000000100000001000000000000000000000005" +
|
||||
"0000000c000000a840e0000041000000411000003ff8000000000000000100020003800140040000" +
|
||||
"000000000004000500068001"
|
||||
raw, err := hex.DecodeString(dump)
|
||||
if err != nil {
|
||||
t.Fatalf("hex: %v", err)
|
||||
}
|
||||
path := filepath.Join(t.TempDir(), "foreign.nc")
|
||||
if err := os.WriteFile(path, raw, 0o644); err != nil {
|
||||
t.Fatalf("write: %v", err)
|
||||
}
|
||||
dims, vars, _, err := LoadNetCDF(path)
|
||||
if err != nil {
|
||||
t.Fatalf("LoadNetCDF: %v", err)
|
||||
}
|
||||
if len(dims) != 2 || dims[0].Name != "time" || dims[0].Length != 0 || dims[1].Length != 3 {
|
||||
t.Fatalf("dims = %+v", dims)
|
||||
}
|
||||
if len(vars) != 3 {
|
||||
t.Fatalf("vars = %d, want 3", len(vars))
|
||||
}
|
||||
byName := map[string]NetCDFVar{}
|
||||
for _, v := range vars {
|
||||
byName[v.Name] = v
|
||||
}
|
||||
// Two records of one double.
|
||||
if got := byName["t"].Values.RawFloats(); len(got) != 2 || got[0] != 1.5 || got[1] != 2.5 {
|
||||
t.Fatalf("t = %v, want [1.5 2.5]", got)
|
||||
}
|
||||
// Two records of three shorts, each slab padded on disk.
|
||||
if got := byName["s"].Values.Shape(); got[0] != 2 || got[1] != 3 {
|
||||
t.Fatalf("s shape = %v, want [2 3]", got)
|
||||
}
|
||||
for i, want := range []float64{1, 2, 3, 4, 5, 6} {
|
||||
if got := byName["s"].Values.FloatAt(i); got != want {
|
||||
t.Fatalf("s[%d] = %v, want %v", i, got, want)
|
||||
}
|
||||
}
|
||||
// The fixed variable, written before the records. Every variable
|
||||
// comes back widened to float64, whatever the file stores.
|
||||
for i, want := range []float64{7, 8, 9} {
|
||||
if got := byName["v"].Values.FloatAt(i); got != want {
|
||||
t.Fatalf("v[%d] = %v, want %v", i, got, want)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// TestNetCDFRecordZeroRecords pins the empty unlimited dimension: a
|
||||
// file may declare the record dimension and hold no records at all.
|
||||
func TestNetCDFRecordZeroRecords(t *testing.T) {
|
||||
path := filepath.Join(t.TempDir(), "empty.nc")
|
||||
empty, err := core.FromFloats([]float64{}, 0)
|
||||
if err != nil {
|
||||
t.Fatalf("FromFloats: %v", err)
|
||||
}
|
||||
err = SaveNetCDF(path, []NetCDFDim{{Name: "time", Length: 0}},
|
||||
[]NetCDFVar{{Name: "t", Dims: []string{"time"}, Values: empty}}, nil)
|
||||
if err != nil {
|
||||
t.Fatalf("SaveNetCDF: %v", err)
|
||||
}
|
||||
dims, vars, _, err := LoadNetCDF(path)
|
||||
if err != nil {
|
||||
t.Fatalf("LoadNetCDF: %v", err)
|
||||
}
|
||||
if len(dims) != 1 || dims[0].Length != 0 {
|
||||
t.Fatalf("dims = %+v", dims)
|
||||
}
|
||||
if len(vars) != 1 || vars[0].Values.Len() != 0 {
|
||||
t.Fatalf("vars = %+v, want one empty variable", vars)
|
||||
}
|
||||
if s := vars[0].Values.Shape(); s[0] != 0 {
|
||||
t.Fatalf("shape = %v, want [0]", s)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,2 @@
|
||||
go test fuzz v1
|
||||
[]byte("\x89HDF\r\n\x1a\n\x000000\b\b0000000000000000000000000000000000000000000000000`\x00\x00\x00\x00\x00\x00\x00000000000000000000000000\x010\x06\x000000\x00\x01\x00\x000000\x01\x00\x18\x000000\x01\x0100000000000\x00\x00\x00\x04\x00\x00\x00\x00\x00\x00\x00\x03\x00\x18\x00\x01\x00\x00\x00\x11 ?\x00\b\x00\x00\x00\x00\x00@\x004\v\x004\xff\x03\x00\x00\x00\x00\x00\x00\x05\x00\b\x00\x01\x00\x00\x00\x02\x03\x02\x01\x00\x00\x00\x00\v\x008\x00\x01\x00\x00\x00\x01\x02\x00\x00\x00\x00\x00\x00\x02\x00\b\x00\x01\x00\x01\x00shuffle\x00\b\x00\x00\x00\x00\x00\x00\x00\x01\x00\b\x00\x01\x00\x01\x00deflate\x00\x04\x00\x00\x00\x00\x00\x00\x00\b\x00\x18\x00\x00\x00\x00\x00\x03\x02\x02\x00 \x00\x00\x00\x00\x7f\x00\x02\x00\x00\x00\b\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00H\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00")
|
||||
@@ -0,0 +1,2 @@
|
||||
go test fuzz v1
|
||||
[]byte("\x89HDF\r\n\x1a\n\x000000\b\b0000000000000000000000000000000000000000000000000`\x00\x00\x00\x00\x00\x00\x00000000000000000000000000\x010\x01\x000000\x18\x00\x00\x000000\x06\x00\a\x000000\x01700000000000000")
|
||||
@@ -0,0 +1,2 @@
|
||||
go test fuzz v1
|
||||
[]byte("\x89HDF\r\n\x1a\n\x000000\b\b0000000000000000000000000000000000000000000000000`\x00\x00\x00\x00\x00\x00\x00000000000000000000000000\x010\x03\x000000\x18\x00\x00\x000000\x06\x00\x00\x00000000\x00\x00000000\x00\x000000")
|
||||
@@ -0,0 +1,2 @@
|
||||
go test fuzz v1
|
||||
[]byte("CDF\x0100000000\x00\x00\x00\x02\x00\x00\x00\x030000\x00\x00\x00\x03\x00\x00\x00\x030000\x00\x00\x00\x040000\x00\x00\x00\x01\x00\x00\x00\x0500000000\x00\x00\x00\x02\x00\x00\x00\x0400000000\x00\x00\x00\x02\x00\x00\x00\x040000\x00\x00\x00\x02\x00\x00\x00\x00\x00\x00\x00\x010000\x00\x00\x00\x00\x00\x00\x00\x060000\x00\x00\x000\x00\x00\x00\x040000\x00\x00\x00\x01\x00\x00\x00\x000000\x00\x00\x00\x00\x00\x00\x00\x040000\x00\x00\x000000000000000000000000000000000000000000000000000000000000000000000000000000000000000000000000000")
|
||||
@@ -0,0 +1,2 @@
|
||||
go test fuzz v1
|
||||
[]byte("CDF\x0100000000\x00\x00\x00\x02\x00\x00\x00\x000000\x00\x00\x00\x01000000000000\x00\x00\x00\x000000\x00\x00\x00\x01\x00\x00\x00\x00\x00\x00\x00\x000000\x00\x00\x00\x00\x00\x00\x00\x010000\x00\x00\x000")
|
||||
Vendored
BIN
Binary file not shown.
Vendored
BIN
Binary file not shown.
Vendored
BIN
Binary file not shown.
Vendored
BIN
Binary file not shown.
@@ -0,0 +1,204 @@
|
||||
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
|
||||
// SPDX-License-Identifier: MIT
|
||||
|
||||
package io
|
||||
|
||||
import (
|
||||
"encoding/binary"
|
||||
"fmt"
|
||||
"math"
|
||||
"runtime"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
)
|
||||
|
||||
// Regression pins for the HDF5 group walk: deep chains that must stay
|
||||
// linear, the depth cap, self-links and hard-link diamonds that must
|
||||
// be refused as cycles.
|
||||
|
||||
// groupChain builds a well-formed HDF5 file holding a chain of
|
||||
// depth groups, each carrying one named attribute and a hard link to the
|
||||
// next. Every group is legal, the attributes are legal scalar attributes
|
||||
// and the chain terminates, so the reader has no grounds to refuse it:
|
||||
// the walk has to stay cheap on its own.
|
||||
func groupChain(depth int) []byte {
|
||||
const (
|
||||
firstGroup = 128
|
||||
stride = 88
|
||||
)
|
||||
n := firstGroup + depth*stride + 16
|
||||
f := make([]byte, n)
|
||||
copy(f, hdf5Magic)
|
||||
f[8] = 0 // superblock version 0
|
||||
f[13] = 8
|
||||
f[14] = 8
|
||||
binary.LittleEndian.PutUint64(f[32:], math.MaxUint64) // free space undefined
|
||||
binary.LittleEndian.PutUint64(f[40:], uint64(n)) // end of file
|
||||
binary.LittleEndian.PutUint64(f[48:], math.MaxUint64) // driver undefined
|
||||
binary.LittleEndian.PutUint64(f[64:], firstGroup) // root object header
|
||||
|
||||
for k := range depth {
|
||||
a := firstGroup + k*stride
|
||||
nmsg := 2
|
||||
if k == depth-1 {
|
||||
nmsg = 1 // the tail group only carries its attribute
|
||||
}
|
||||
f[a] = 1 // object header version 1
|
||||
binary.LittleEndian.PutUint16(f[a+2:], uint16(nmsg))
|
||||
binary.LittleEndian.PutUint32(f[a+4:], 1) // reference count
|
||||
binary.LittleEndian.PutUint32(f[a+8:], 88) // message data size
|
||||
|
||||
// Attribute message (type 12), size 40, body at a+24.
|
||||
binary.LittleEndian.PutUint16(f[a+16:], hdf5MsgAttribute)
|
||||
binary.LittleEndian.PutUint16(f[a+18:], 40)
|
||||
f[a+24] = 1 // version
|
||||
binary.LittleEndian.PutUint16(f[a+26:], 8) // name length
|
||||
binary.LittleEndian.PutUint16(f[a+28:], 8) // datatype length
|
||||
binary.LittleEndian.PutUint16(f[a+30:], 8) // dataspace length
|
||||
copy(f[a+32:], fmt.Sprintf("a%07d", k)) // 8-byte attribute name
|
||||
f[a+40] = 0x10 // fixed-point datatype
|
||||
binary.LittleEndian.PutUint32(f[a+44:], 1) // one byte
|
||||
f[a+48] = 1 // scalar dataspace
|
||||
f[a+56] = byte(k) // one byte of value
|
||||
|
||||
if k == depth-1 {
|
||||
continue
|
||||
}
|
||||
// Link message (type 6), size 12, body at a+72.
|
||||
binary.LittleEndian.PutUint16(f[a+64:], hdf5MsgLink)
|
||||
binary.LittleEndian.PutUint16(f[a+66:], 12)
|
||||
f[a+72] = 1 // version
|
||||
f[a+73] = 0 // flags: hard link, 1-byte name length
|
||||
f[a+74] = 1 // name length
|
||||
f[a+75] = 'a'
|
||||
binary.LittleEndian.PutUint64(f[a+76:], uint64(a+stride))
|
||||
}
|
||||
return f
|
||||
}
|
||||
|
||||
// selfLink builds the smallest HDF5 file whose root group links to
|
||||
// itself: 136 bytes, one object header, one link message.
|
||||
func selfLink() []byte {
|
||||
const N = 136
|
||||
f := make([]byte, N)
|
||||
copy(f, hdf5Magic)
|
||||
f[8] = 0 // superblock version 0
|
||||
f[13] = 8
|
||||
f[14] = 8
|
||||
binary.LittleEndian.PutUint64(f[32:], math.MaxUint64)
|
||||
binary.LittleEndian.PutUint64(f[40:], N)
|
||||
binary.LittleEndian.PutUint64(f[48:], math.MaxUint64)
|
||||
binary.LittleEndian.PutUint64(f[64:], 96) // root object header
|
||||
|
||||
f[96] = 1 // object header version 1
|
||||
binary.LittleEndian.PutUint16(f[98:], 1)
|
||||
binary.LittleEndian.PutUint32(f[100:], 1)
|
||||
binary.LittleEndian.PutUint32(f[104:], 24)
|
||||
|
||||
binary.LittleEndian.PutUint16(f[112:], hdf5MsgLink)
|
||||
binary.LittleEndian.PutUint16(f[114:], 12)
|
||||
f[120] = 1 // link message version
|
||||
f[121] = 0 // hard link, 1-byte name length
|
||||
f[122] = 1 // name length
|
||||
f[123] = 'a'
|
||||
binary.LittleEndian.PutUint64(f[124:], 96) // the link target: itself
|
||||
return f
|
||||
}
|
||||
|
||||
// TestWalkRefusesHardLinkCycle pins the crash that took the editor down:
|
||||
// walk had no visited set, so a group linked into itself recursed for
|
||||
// ever while the path string grew, and the heap grew with it at hundreds
|
||||
// of megabytes per second until the host ran out of memory. The reader
|
||||
// must refuse the file with an error.
|
||||
func TestWalkRefusesHardLinkCycle(t *testing.T) {
|
||||
guard := time.AfterFunc(20*time.Second, func() { panic("LoadHDF5 on a cyclic file did not return") })
|
||||
defer guard.Stop()
|
||||
|
||||
path := writeHostile(t, "selflink.h5", selfLink())
|
||||
stop := warnHeap(t, 256<<20)
|
||||
_, err := LoadHDF5(path)
|
||||
stop()
|
||||
if err == nil {
|
||||
t.Fatal("LoadHDF5 accepted a group that hard-links into itself")
|
||||
}
|
||||
if !strings.Contains(err.Error(), "hard-link cycle") {
|
||||
t.Fatalf("LoadHDF5 = %v, want the cycle refusal", err)
|
||||
}
|
||||
}
|
||||
|
||||
// TestWalkDeepChainStaysLinear pins the other half of the same crash: a
|
||||
// well-formed chain of groups used to cost memory quadratic in its
|
||||
// depth, because every level copied the inherited attribute map and
|
||||
// built a longer path string. The live memory must grow with the depth,
|
||||
// not with its square, and the walk carries one shared path buffer and
|
||||
// one shared attribute map to make that so.
|
||||
func TestWalkDeepChainStaysLinear(t *testing.T) {
|
||||
guard := time.AfterFunc(60*time.Second, func() { panic("deep chain walk did not return") })
|
||||
defer guard.Stop()
|
||||
|
||||
measure := func(depth int) (uint64, int) {
|
||||
data := groupChain(depth)
|
||||
path := writeHostile(t, fmt.Sprintf("chain%d.h5", depth), data)
|
||||
runtime.GC()
|
||||
var before, after runtime.MemStats
|
||||
runtime.ReadMemStats(&before)
|
||||
if _, err := LoadHDF5(path); err != nil {
|
||||
t.Fatalf("depth %d: LoadHDF5 refused a well-formed file: %v", depth, err)
|
||||
}
|
||||
runtime.ReadMemStats(&after)
|
||||
return after.TotalAlloc - before.TotalAlloc, len(data)
|
||||
}
|
||||
small, smallFile := measure(150)
|
||||
big, bigFile := measure(450)
|
||||
// Linear: three times the depth is about three times the bytes. The
|
||||
// per-level copies this replaced measured 16.5x for the same step.
|
||||
if ratio := float64(big) / float64(small); ratio > 6 {
|
||||
t.Errorf("allocation grows with the square of the depth: %.1fx for 3x depth (%d B for %d B, %d B for %d B)",
|
||||
ratio, big, bigFile, small, smallFile)
|
||||
}
|
||||
}
|
||||
|
||||
// TestWalkRefusesUnboundedNesting pins the depth cap: past it the reader
|
||||
// refuses the file loudly and cheaply instead of walking every level
|
||||
// first.
|
||||
func TestWalkRefusesUnboundedNesting(t *testing.T) {
|
||||
guard := time.AfterFunc(60*time.Second, func() { panic("deep chain walk did not return") })
|
||||
defer guard.Stop()
|
||||
|
||||
deep := writeHostile(t, "deep.h5", groupChain(2*hdf5MaxGroupDepth))
|
||||
stop := warnHeap(t, 256<<20)
|
||||
_, err := LoadHDF5(deep)
|
||||
stop()
|
||||
if err == nil {
|
||||
t.Fatalf("LoadHDF5 accepted a chain %d groups deep", 2*hdf5MaxGroupDepth)
|
||||
}
|
||||
if !strings.Contains(err.Error(), "nest deeper than") {
|
||||
t.Fatalf("LoadHDF5 = %v, want the nesting refusal", err)
|
||||
}
|
||||
}
|
||||
|
||||
// warnHeap fails the test as soon as the live heap passes limit, so a
|
||||
// regression cannot consume the machine: the watchdog panics, which
|
||||
// unwinds the offending walk instead of letting it allocate on.
|
||||
func warnHeap(t *testing.T, limit uint64) func() {
|
||||
t.Helper()
|
||||
done := make(chan struct{})
|
||||
go func() {
|
||||
ticker := time.NewTicker(20 * time.Millisecond)
|
||||
defer ticker.Stop()
|
||||
for {
|
||||
select {
|
||||
case <-done:
|
||||
return
|
||||
case <-ticker.C:
|
||||
var m runtime.MemStats
|
||||
runtime.ReadMemStats(&m)
|
||||
if m.HeapAlloc > limit {
|
||||
panic("walk heap above its cap")
|
||||
}
|
||||
}
|
||||
}
|
||||
}()
|
||||
return func() { close(done) }
|
||||
}
|
||||
@@ -0,0 +1,61 @@
|
||||
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
|
||||
// SPDX-License-Identifier: MIT
|
||||
|
||||
package io
|
||||
|
||||
// Benchmarks for the paths the coordinator's battery does not isolate.
|
||||
// Every fixture is deterministic and built before the measured loop.
|
||||
|
||||
import (
|
||||
"os"
|
||||
"path/filepath"
|
||||
"testing"
|
||||
)
|
||||
|
||||
// BenchmarkWaveFITSBinaryTextCells measures the character-column decode
|
||||
// alone: one 8A column of four thousand rows, the per-cell work the
|
||||
// binary-table benchmark mixes with seven numeric columns.
|
||||
func BenchmarkWaveFITSBinaryTextCells(b *testing.B) {
|
||||
const rows = 4000
|
||||
path := filepath.Join(b.TempDir(), "textcells.fits")
|
||||
body := make([]byte, rows*8)
|
||||
for r := range rows {
|
||||
copy(body[r*8:], "star ")
|
||||
}
|
||||
cards := []string{
|
||||
fitsStringCardRaw("XTENSION", "BINTABLE"),
|
||||
fitsIntCard("BITPIX", 8),
|
||||
fitsIntCard("NAXIS", 2),
|
||||
fitsIntCard("NAXIS1", 8),
|
||||
fitsIntCard("NAXIS2", rows),
|
||||
fitsIntCard("PCOUNT", 0),
|
||||
fitsIntCard("GCOUNT", 1),
|
||||
fitsIntCard("TFIELDS", 1),
|
||||
fitsStringCardRaw("TTYPE1", "STAR"),
|
||||
fitsStringCardRaw("TFORM1", "8A"),
|
||||
fitsEndCard(),
|
||||
}
|
||||
out := fitsAppendCards(nil, []string{
|
||||
fitsBoolCard("SIMPLE", true),
|
||||
fitsIntCard("BITPIX", 8),
|
||||
fitsIntCard("NAXIS", 0),
|
||||
fitsBoolCard("EXTEND", true),
|
||||
fitsEndCard(),
|
||||
})
|
||||
out = fitsAppendCards(out, cards)
|
||||
out = append(out, body...)
|
||||
out = fitsAppendZeroPad(out)
|
||||
if err := os.WriteFile(path, out, 0o644); err != nil {
|
||||
b.Fatal(err)
|
||||
}
|
||||
b.ReportAllocs()
|
||||
for b.Loop() {
|
||||
table, err := LoadFITSTable(path)
|
||||
if err != nil {
|
||||
b.Fatal(err)
|
||||
}
|
||||
if table.Rows != rows || table.Text[0] == nil || table.Text[0][0] != "star" {
|
||||
b.Fatalf("table came back as %s with %d rows", table.Kind, table.Rows)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,61 @@
|
||||
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
|
||||
// SPDX-License-Identifier: MIT
|
||||
|
||||
package io
|
||||
|
||||
import (
|
||||
"flag"
|
||||
"fmt"
|
||||
"os"
|
||||
"runtime"
|
||||
"testing"
|
||||
"time"
|
||||
)
|
||||
|
||||
// TestMain installs a heap watchdog for the whole package run. The io
|
||||
// tests are hostile-input tests by design: they feed crafted headers to
|
||||
// the readers and check that nothing panics and nothing grows without
|
||||
// bound. A regression therefore cannot fail politely, it allocates: a
|
||||
// hard-link cycle and a deep group chain both took the host down once.
|
||||
// The watchdog turns any such runaway into an immediate panic that
|
||||
// unwinds the offending code path, so the worst case is a failed test
|
||||
// rather than an OOM kill of the editor that started it.
|
||||
func TestMain(m *testing.M) {
|
||||
// The aggregate read budget bounds what one input may allocate
|
||||
// legally; the watchdog sits an order above it. Fuzzing runs many
|
||||
// workers holding inputs at once and the engine carries its own
|
||||
// corpus and coverage bookkeeping, so the fuzz run watches a wider
|
||||
// ceiling while still catching the unbounded growth the guard
|
||||
// exists for: a runaway reaches any finite cap in seconds.
|
||||
//
|
||||
// The flags are parsed here, before the ceiling is decided: the
|
||||
// test flags are only registered at TestMain and m.Run parses them
|
||||
// later, so an earlier read of test.fuzz sees its empty default and
|
||||
// the fuzz ceiling never fires.
|
||||
flag.Parse()
|
||||
const plainCap = 3 << 30
|
||||
ceiling := int64(plainCap)
|
||||
if f := flag.Lookup("test.fuzz"); f != nil && f.Value.String() != "" {
|
||||
ceiling = 32 << 30
|
||||
}
|
||||
stop := make(chan struct{})
|
||||
go func() {
|
||||
ticker := time.NewTicker(20 * time.Millisecond)
|
||||
defer ticker.Stop()
|
||||
for {
|
||||
select {
|
||||
case <-stop:
|
||||
return
|
||||
case <-ticker.C:
|
||||
var m runtime.MemStats
|
||||
runtime.ReadMemStats(&m)
|
||||
if int64(m.HeapAlloc) > ceiling {
|
||||
panic(fmt.Sprintf("io test binary heap above %d GiB: a reader is allocating without bound", ceiling>>30))
|
||||
}
|
||||
}
|
||||
}
|
||||
}()
|
||||
code := m.Run()
|
||||
close(stop)
|
||||
os.Exit(code)
|
||||
}
|
||||
Reference in New Issue
Block a user