From f705a517ae64808fa5c4b6809d46a6f5c298a6fb Mon Sep 17 00:00:00 2001 From: apstndb <803393+apstndb@users.noreply.github.com> Date: Fri, 10 Jul 2026 04:11:06 +0900 Subject: [PATCH 1/2] fix: reject extra parameter file documents --- params/load.go | 23 +++++++++++- params/load_test.go | 89 +++++++++++++++++++++++++++++++++++++++++++++ 2 files changed, 111 insertions(+), 1 deletion(-) diff --git a/params/load.go b/params/load.go index e1f282a..4518959 100644 --- a/params/load.go +++ b/params/load.go @@ -3,7 +3,9 @@ package params import ( "bytes" "encoding/json" + "errors" "fmt" + "io" "maps" "math" "os" @@ -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) } @@ -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 { return nil, fmt.Errorf("parse param file as YAML: %w", err) } + if raw == nil { + return nil, fmt.Errorf("parse param file as YAML: 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 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) diff --git a/params/load_test.go b/params/load_test.go index 55477c0..741eae0 100644 --- a/params/load_test.go +++ b/params/load_test.go @@ -109,6 +109,95 @@ func TestLoadParamFile(t *testing.T) { } } +func TestLoadParamFileAcceptsSingleMappingDocument(t *testing.T) { + t.Parallel() + + tests := []struct { + name string + filename string + content string + }{ + { + name: "JSON with trailing whitespace", + filename: "params.json", + content: "{} \n\t", + }, + { + name: "YAML empty mapping", + filename: "params.yaml", + content: "{}\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) + } + got, err := LoadParamFile(path) + if err != nil { + t.Fatal(err) + } + if got == nil || len(got) != 0 { + t.Fatalf("expected an empty map, got %v", got) + } + }) + } +} + +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() From e71a81ac39f78ecd6886d43bf10ebcfc0c10abf3 Mon Sep 17 00:00:00 2001 From: apstndb <803393+apstndb@users.noreply.github.com> Date: Fri, 10 Jul 2026 04:18:22 +0900 Subject: [PATCH 2/2] Preserve empty YAML parameter files --- params/load.go | 6 +++--- params/load_test.go | 25 +++++++++++++++++++++++-- 2 files changed, 26 insertions(+), 5 deletions(-) diff --git a/params/load.go b/params/load.go index 4518959..d0a761b 100644 --- a/params/load.go +++ b/params/load.go @@ -100,11 +100,11 @@ func loadParamYAMLFile(b []byte) (map[string]string, error) { var raw map[string]yaml.RawMessage 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 raw == nil { - return nil, fmt.Errorf("parse param file as YAML: 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 YAML: %w", err) diff --git a/params/load_test.go b/params/load_test.go index 741eae0..7e7019a 100644 --- a/params/load_test.go +++ b/params/load_test.go @@ -116,16 +116,37 @@ func TestLoadParamFileAcceptsSingleMappingDocument(t *testing.T) { 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"}, }, } @@ -141,8 +162,8 @@ func TestLoadParamFileAcceptsSingleMappingDocument(t *testing.T) { if err != nil { t.Fatal(err) } - if got == nil || len(got) != 0 { - t.Fatalf("expected an empty map, got %v", got) + if diff := cmp.Diff(tc.want, got); diff != "" { + t.Fatalf("LoadParamFile() mismatch (-want +got):\n%s", diff) } }) }