From 5020b688d79ebb5eb81a6623d65e67db6a5475a7 Mon Sep 17 00:00:00 2001 From: AdamMagued Date: Fri, 2 Oct 2026 12:45:42 +0000 Subject: [PATCH] fix(decode): reject requests with trailing non-whitespace data in DecodeJSON json.Decoder.Decode reads a single JSON value and leaves the decoder stream positioned after that value. Previously, DecodeJSON returned immediately after the first Decode call without verifying that the stream had ended, causing requests with trailing non-whitespace data or extra JSON values after the valid payload to be accepted. Ensure no trailing non-whitespace data exists after decoding the target value by checking for io.EOF on a subsequent decode call. Add unit and regression tests for DecodeJSON and DefaultDecoder. Fixes #42 --- decoder.go | 14 +++++- decoder_test.go | 127 ++++++++++++++++++++++++++++++++++++++++++++++++ 2 files changed, 140 insertions(+), 1 deletion(-) create mode 100644 decoder_test.go diff --git a/decoder.go b/decoder.go index c01be83..85c826b 100644 --- a/decoder.go +++ b/decoder.go @@ -40,7 +40,19 @@ func DefaultDecoder(r *http.Request, v interface{}) error { // DecodeJSON decodes a given reader into an interface using the json decoder. func DecodeJSON(r io.Reader, v interface{}) error { defer io.Copy(io.Discard, r) //nolint:errcheck - return json.NewDecoder(r).Decode(v) + + dec := json.NewDecoder(r) + if err := dec.Decode(v); err != nil { + return err + } + var extra struct{} + if err := dec.Decode(&extra); err != io.EOF { + if err == nil { + return errors.New("render: extra data after JSON payload") + } + return err + } + return nil } // DecodeXML decodes a given reader into an interface using the xml decoder. diff --git a/decoder_test.go b/decoder_test.go new file mode 100644 index 0000000..c01842e --- /dev/null +++ b/decoder_test.go @@ -0,0 +1,127 @@ +package render + +import ( + "bytes" + "net/http" + "net/http/httptest" + "strings" + "testing" +) + +func TestDecodeJSON(t *testing.T) { + t.Run("valid json object", func(t *testing.T) { + var v struct { + Title string `json:"title"` + } + err := DecodeJSON(strings.NewReader(`{"title": "hello"}`), &v) + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + if v.Title != "hello" { + t.Fatalf("expected hello, got %s", v.Title) + } + }) + + t.Run("valid json with trailing whitespace", func(t *testing.T) { + var v map[string]interface{} + err := DecodeJSON(strings.NewReader("{\"title\": \"hello\"} \n\t\r\n "), &v) + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + if v["title"] != "hello" { + t.Fatalf("expected hello, got %v", v["title"]) + } + }) + + t.Run("rejects trailing garbage data", func(t *testing.T) { + var v map[string]interface{} + // Regression test for issue #42 + err := DecodeJSON(strings.NewReader("{}\n some garbage data"), &v) + if err == nil { + t.Fatal("expected error when trailing garbage data follows valid JSON, got nil") + } + }) + + t.Run("rejects trailing second json object", func(t *testing.T) { + var v map[string]interface{} + err := DecodeJSON(strings.NewReader(`{"a": 1} {"b": 2}`), &v) + if err == nil { + t.Fatal("expected error when second JSON object follows, got nil") + } + }) + + t.Run("rejects trailing primitive", func(t *testing.T) { + var v map[string]interface{} + err := DecodeJSON(strings.NewReader(`{"a": 1} 123`), &v) + if err == nil { + t.Fatal("expected error when trailing number follows JSON, got nil") + } + }) + + t.Run("rejects trailing delimiter", func(t *testing.T) { + var v map[string]interface{} + err := DecodeJSON(strings.NewReader(`{"a": 1} ]`), &v) + if err == nil { + t.Fatal("expected error when trailing delimiter follows JSON, got nil") + } + + err = DecodeJSON(strings.NewReader(`{"a": 1} }`), &v) + if err == nil { + t.Fatal("expected error when trailing delimiter follows JSON, got nil") + } + }) + + t.Run("rejects trailing null", func(t *testing.T) { + var v map[string]interface{} + err := DecodeJSON(strings.NewReader(`{"a": 1} null`), &v) + if err == nil { + t.Fatal("expected error when trailing null follows JSON, got nil") + } + }) + + t.Run("invalid json returns error", func(t *testing.T) { + var v map[string]interface{} + err := DecodeJSON(strings.NewReader(`{invalid`), &v) + if err == nil { + t.Fatal("expected error on invalid JSON, got nil") + } + }) +} + +func TestDefaultDecoder(t *testing.T) { + t.Run("json with trailing garbage returns error", func(t *testing.T) { + req := httptest.NewRequest(http.MethodPost, "/", strings.NewReader("{}\n some garbage data")) + req.Header.Set("Content-Type", "application/json") + + var v map[string]interface{} + err := DefaultDecoder(req, &v) + if err == nil { + t.Fatal("expected error for request with trailing garbage, got nil") + } + }) + + t.Run("valid json succeeds", func(t *testing.T) { + req := httptest.NewRequest(http.MethodPost, "/", bytes.NewBufferString(`{"status": "ok"}`)) + req.Header.Set("Content-Type", "application/json") + + var v map[string]string + err := DefaultDecoder(req, &v) + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + if v["status"] != "ok" { + t.Fatalf("expected ok, got %v", v["status"]) + } + }) + + t.Run("unsupported content type returns error", func(t *testing.T) { + req := httptest.NewRequest(http.MethodPost, "/", strings.NewReader("plain text")) + req.Header.Set("Content-Type", "text/plain") + + var v map[string]interface{} + err := DefaultDecoder(req, &v) + if err == nil { + t.Fatal("expected error for unsupported content type, got nil") + } + }) +}