Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
23 changes: 22 additions & 1 deletion params/load.go
Original file line number Diff line number Diff line change
Expand Up @@ -3,7 +3,9 @@ package params
import (
"bytes"
"encoding/json"
"errors"
"fmt"
"io"
"maps"
"math"
"os"
Expand Down Expand Up @@ -71,6 +73,15 @@ func LoadParamFile(path string) (map[string]string, error) {
if err := dec.Decode(&raw); err != nil {
return nil, fmt.Errorf("parse param file as JSON: %w", err)
}
if raw == nil {
return nil, fmt.Errorf("parse param file as JSON: top-level value must be a mapping")
}
if err := dec.Decode(new(any)); !errors.Is(err, io.EOF) {
if err != nil {
return nil, fmt.Errorf("parse param file as JSON: %w", err)
}
return nil, fmt.Errorf("parse param file as JSON: multiple top-level values")
}
default:
return loadParamYAMLFile(b)
}
Expand All @@ -87,9 +98,19 @@ func LoadParamFile(path string) (map[string]string, error) {

func loadParamYAMLFile(b []byte) (map[string]string, error) {
var raw map[string]yaml.RawMessage
if err := yaml.Unmarshal(b, &raw); err != nil {
dec := yaml.NewDecoder(bytes.NewReader(b))
if err := dec.Decode(&raw); err != nil {
if errors.Is(err, io.EOF) {
return map[string]string{}, nil
}
return nil, fmt.Errorf("parse param file as YAML: %w", err)
}
if err := dec.Decode(new(any)); !errors.Is(err, io.EOF) {
if err != nil {
return nil, fmt.Errorf("parse param file as YAML: %w", err)
}
return nil, fmt.Errorf("parse param file as YAML: multiple documents are not supported")
}
out := make(map[string]string, len(raw))
for k, msg := range raw {
s, err := paramFileYAMLValueToString(msg)
Expand Down
110 changes: 110 additions & 0 deletions params/load_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -109,6 +109,116 @@ func TestLoadParamFile(t *testing.T) {
}
}

func TestLoadParamFileAcceptsSingleMappingDocument(t *testing.T) {
t.Parallel()

tests := []struct {
name string
filename string
content string
want map[string]string
}{
{
name: "JSON with trailing whitespace",
filename: "params.json",
content: "{} \n\t",
want: map[string]string{},
},
{
name: "YAML empty mapping",
filename: "params.yaml",
content: "{}\n",
want: map[string]string{},
},
{
name: "YAML comments only",
filename: "params.yaml",
content: "# no parameters yet\n",
want: map[string]string{},
},
{
name: "YAML empty document marker",
filename: "params.yaml",
content: "---\n",
want: map[string]string{},
},
{
name: "YAML trailing empty document marker",
filename: "params.yaml",
content: "x: INT64\n---\n",
want: map[string]string{"x": "INT64"},
},
}

for _, tc := range tests {
t.Run(tc.name, func(t *testing.T) {
t.Parallel()

path := filepath.Join(t.TempDir(), tc.filename)
if err := os.WriteFile(path, []byte(tc.content), 0o644); err != nil {
t.Fatal(err)
}
got, err := LoadParamFile(path)
if err != nil {
t.Fatal(err)
}
if diff := cmp.Diff(tc.want, got); diff != "" {
t.Fatalf("LoadParamFile() mismatch (-want +got):\n%s", diff)
}
})
}
}

func TestLoadParamFileRejectsInvalidDocuments(t *testing.T) {
t.Parallel()

tests := []struct {
name string
filename string
content string
}{
{
name: "JSON trailing non-whitespace content",
filename: "params.json",
content: `{"x":1} trailing`,
},
{
name: "JSON second value",
filename: "params.json",
content: `{"x":1} {"y":2}`,
},
{
name: "JSON top-level null",
filename: "params.json",
content: `null`,
},
{
name: "YAML second document",
filename: "params.yaml",
content: "x: 1\n---\ny: 2\n",
},
{
name: "YAML top-level null",
filename: "params.yaml",
content: "null\n",
},
}

for _, tc := range tests {
t.Run(tc.name, func(t *testing.T) {
t.Parallel()

path := filepath.Join(t.TempDir(), tc.filename)
if err := os.WriteFile(path, []byte(tc.content), 0o644); err != nil {
t.Fatal(err)
}
if _, err := LoadParamFile(path); err == nil {
t.Fatal("expected an error")
}
})
}
}

func TestFormatParamFloat(t *testing.T) {
t.Parallel()

Expand Down
Loading