util: decode files by their extension

UnmarshalByExt picks JSON, HuJSON or YAML from the file name; content can't
tell them apart, since YAML parses JSON syntax.
This commit is contained in:
Kristoffer Dalby
2026-09-28 13:31:51 +00:00
committed by Kristoffer Dalby
parent 90732bdaaf
commit d20a412079
2 changed files with 107 additions and 0 deletions
+57
View File
@@ -1,6 +1,7 @@
package util
import (
"errors"
"net/netip"
"strings"
"testing"
@@ -891,3 +892,59 @@ func TestGenerateRegistrationKey(t *testing.T) {
})
}
}
func TestUnmarshalByExt(t *testing.T) {
type rec struct {
Name string `json:"name"`
Port int `json:"port"`
}
want := rec{Name: "a", Port: 1}
tests := []struct {
name string
file string
data string
wantFormat bool // want ErrUnknownFileFormat
wantErr bool
}{
{name: "json", file: "x.json", data: `{"name": "a", "port": 1}`},
{name: "extension case is ignored", file: "x.JSON", data: `{"name": "a", "port": 1}`},
{name: "hujson", file: "x.hujson", data: "{\n // comment\n \"name\": \"a\",\n \"port\": 1,\n}"},
{name: "yaml", file: "x.yaml", data: "name: a\nport: 1\n"},
{name: "yml", file: "x.yml", data: "name: a\nport: 1\n"},
// The extension is binding: a .json file does not get HuJSON leniency.
{name: "comments in json", file: "x.json", data: "{\n // comment\n \"name\": \"a\"\n}", wantErr: true},
{name: "unknown extension", file: "x.txt", data: `{"name": "a", "port": 1}`, wantFormat: true},
{name: "no extension", file: "x", data: `{"name": "a", "port": 1}`, wantFormat: true},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
got, err := UnmarshalByExt[rec](tt.file, []byte(tt.data))
if tt.wantFormat {
if !errors.Is(err, ErrUnknownFileFormat) {
t.Fatalf("UnmarshalByExt(%q) error = %v, want %v", tt.file, err, ErrUnknownFileFormat)
}
return
}
if tt.wantErr {
if err == nil {
t.Fatalf("UnmarshalByExt(%q) error = nil, want error", tt.file)
}
return
}
if err != nil {
t.Fatalf("UnmarshalByExt(%q) error = %v", tt.file, err)
}
if diff := cmp.Diff(want, got); diff != "" {
t.Errorf("UnmarshalByExt(%q) mismatch (-want +got):\n%s", tt.file, diff)
}
})
}
}