diff --git a/surf/router_test.go b/surf/router_test.go index fa5b118..7a7e6ef 100644 --- a/surf/router_test.go +++ b/surf/router_test.go @@ -92,6 +92,120 @@ func TestRawGroupPanicBare500(t *testing.T) { } } +func TestRecoverDiscardsPartialResponse(t *testing.T) { + const secretBody = "secret-partial" + + t.Run("house", func(t *testing.T) { + r := New(nil) + r.Get("/panic", func(w http.ResponseWriter, _ *http.Request) { + w.Header().Set("X-Partial", "secret") + w.WriteHeader(http.StatusAccepted) + _, _ = w.Write([]byte(secretBody)) + panic("secret internals") + }) + h, err := r.compile() + if err != nil { + t.Fatal(err) + } + + rec := httptest.NewRecorder() + h.ServeHTTP(rec, httptest.NewRequest(http.MethodGet, "/panic", nil)) + if rec.Code != http.StatusInternalServerError { + t.Fatalf("status = %d", rec.Code) + } + if got := rec.Header().Get("Content-Type"); got != "application/json" { + t.Fatalf("Content-Type = %q", got) + } + if got := rec.Header().Get("X-Partial"); got != "" { + t.Fatalf("X-Partial leaked: %q", got) + } + const want = `{"error":true,"message":"Internal server error"}` + if got := rec.Body.String(); got != want { + t.Fatalf("body = %q, want %q", got, want) + } + if strings.Contains(rec.Body.String(), secretBody) { + t.Fatalf("partial body leaked: %q", rec.Body.String()) + } + }) + + t.Run("raw", func(t *testing.T) { + r := New(nil) + r.GroupRaw("/oauth", nil, func(g pact.Router) { + g.Get("/panic", func(w http.ResponseWriter, _ *http.Request) { + w.Header().Set("X-Partial", "secret") + w.WriteHeader(http.StatusAccepted) + _, _ = w.Write([]byte(secretBody)) + panic("secret internals") + }) + }) + h, err := r.compile() + if err != nil { + t.Fatal(err) + } + + rec := httptest.NewRecorder() + h.ServeHTTP(rec, httptest.NewRequest(http.MethodGet, "/oauth/panic", nil)) + if rec.Code != http.StatusInternalServerError { + t.Fatalf("status = %d", rec.Code) + } + if rec.Body.Len() != 0 { + t.Fatalf("body = %q", rec.Body.String()) + } + if got := rec.Header().Get("Content-Type"); got != "" { + t.Fatalf("Content-Type = %q", got) + } + if got := rec.Header().Get("X-Partial"); got != "" { + t.Fatalf("X-Partial leaked: %q", got) + } + if strings.Contains(rec.Body.String(), secretBody) { + t.Fatalf("partial body leaked: %q", rec.Body.String()) + } + }) +} + +func TestBufferedResponseCommitsSuccessfulOutput(t *testing.T) { + r := New(nil) + r.Get("/explicit", func(w http.ResponseWriter, _ *http.Request) { + w.Header().Add("X-Result", "one") + w.Header().Add("X-Result", "two") + w.WriteHeader(http.StatusCreated) + w.WriteHeader(http.StatusTeapot) + _, _ = w.Write([]byte("created")) + }) + r.Get("/implicit", func(w http.ResponseWriter, _ *http.Request) { + w.Header().Set("X-Result", "implicit") + _, _ = w.Write([]byte("ok")) + }) + h, err := r.compile() + if err != nil { + t.Fatal(err) + } + + rec := httptest.NewRecorder() + h.ServeHTTP(rec, httptest.NewRequest(http.MethodGet, "/explicit", nil)) + if rec.Code != http.StatusCreated { + t.Fatalf("explicit status = %d", rec.Code) + } + if got := rec.Header().Values("X-Result"); len(got) != 2 || got[0] != "one" || got[1] != "two" { + t.Fatalf("explicit X-Result = %q", got) + } + if got := rec.Body.String(); got != "created" { + t.Fatalf("explicit body = %q", got) + } + + rec = httptest.NewRecorder() + h.ServeHTTP(rec, httptest.NewRequest(http.MethodGet, "/implicit", nil)) + if rec.Code != http.StatusOK { + t.Fatalf("implicit status = %d", rec.Code) + } + if got := rec.Header().Get("X-Result"); got != "implicit" { + t.Fatalf("implicit X-Result = %q", got) + } + if got := rec.Body.String(); got != "ok" { + t.Fatalf("implicit body = %q", got) + } +} + func TestHouseMiddlewareDuplicateNameFailsBoot(t *testing.T) { identity := func(next http.Handler) http.Handler { return next } named := assemblePlugin{