From d20a412079b58b71f6084e65b49cdf99be2c4ea3 Mon Sep 17 00:00:00 2001 From: Kristoffer Dalby Date: Mon, 28 Sep 2026 13:31:51 +0000 Subject: [PATCH] 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. --- hscontrol/util/file.go | 50 ++++++++++++++++++++++++++++++++ hscontrol/util/util_test.go | 57 +++++++++++++++++++++++++++++++++++++ 2 files changed, 107 insertions(+) diff --git a/hscontrol/util/file.go b/hscontrol/util/file.go index f6b09838..a9bc2efe 100644 --- a/hscontrol/util/file.go +++ b/hscontrol/util/file.go @@ -1,6 +1,7 @@ package util import ( + "encoding/json" "errors" "fmt" "io/fs" @@ -10,6 +11,8 @@ import ( "strings" "github.com/spf13/viper" + "github.com/tailscale/hujson" + "gopkg.in/yaml.v3" ) const ( @@ -24,6 +27,53 @@ const ( // ErrDirectoryPermission is returned when creating a directory fails due to permission issues. var ErrDirectoryPermission = errors.New("creating directory failed with permission error") +// ErrUnknownFileFormat is returned for a file whose extension names no format +// [UnmarshalByExt] reads. +var ErrUnknownFileFormat = errors.New("unknown file format, want .json, .hujson, .yaml or .yml") + +// UnmarshalByExt decodes data into a T in the format name's extension picks: +// .json, .hujson (JSON with comments and trailing commas), or .yaml/.yml. +// The extension decides, not the content: YAML parses JSON syntax, so sniffing +// cannot tell the two apart. YAML keys are the lowercased Go field names. +func UnmarshalByExt[T any](name string, data []byte) (T, error) { + var ( + v T + err error + ) + + switch ext := strings.ToLower(filepath.Ext(name)); ext { + case ".json": + err = json.Unmarshal(data, &v) + case ".hujson": + data, err = hujson.Standardize(data) + if err == nil { + err = json.Unmarshal(data, &v) + } + case ".yaml", ".yml": + err = yaml.Unmarshal(data, &v) + default: + return v, fmt.Errorf("%s: %w", name, ErrUnknownFileFormat) + } + + if err != nil { + return v, fmt.Errorf("decoding %s: %w", name, err) + } + + return v, nil +} + +// ReadFileByExt reads the file at path and decodes it with [UnmarshalByExt]. +func ReadFileByExt[T any](path string) (T, error) { + data, err := os.ReadFile(path) + if err != nil { + var zero T + + return zero, err + } + + return UnmarshalByExt[T](path, data) +} + func AbsolutePathFromConfigPath(path string) string { // If a relative path is provided, prefix it with the directory where // the config file was found. diff --git a/hscontrol/util/util_test.go b/hscontrol/util/util_test.go index 656c5b36..80824d5e 100644 --- a/hscontrol/util/util_test.go +++ b/hscontrol/util/util_test.go @@ -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) + } + }) + } +}