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