mirror of
https://github.com/juanfont/headscale.git
synced 2026-09-30 20:09:36 +09:00
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:
committed by
Kristoffer Dalby
parent
90732bdaaf
commit
d20a412079
@@ -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.
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user