package chimw import ( "io" "net/http" "net/http/httptest" "strconv" "strings" "testing" "github.com/go-chi/chi/v5" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" "sourcecraft.dev/bigbes/sr-ht-ecore/pages" ) // refusal is what a renderer was asked for, so a test can assert on the call // rather than on prose rendered by a template this package does not own. type refusal struct { status int message string method string } // recordRefusals returns an ErrorRenderer of the donors' shape, plus the slice // it appends every call to. func recordRefusals() (ErrorRenderer, *[]refusal) { var calls []refusal render := func(w http.ResponseWriter, r *http.Request, status int, message string) { calls = append(calls, refusal{status: status, message: message, method: r.Method}) w.WriteHeader(status) _, _ = io.WriteString(w, message) } return render, &calls } func TestGetHeadServesAHeadThroughTheGetHandler(t *testing.T) { r := chi.NewRouter() var seen []string GetHead(r, "/page", func(w http.ResponseWriter, r *http.Request) { seen = append(seen, r.Method) _, _ = io.WriteString(w, "body") }) for _, method := range []string{http.MethodGet, http.MethodHead} { rec := httptest.NewRecorder() r.ServeHTTP(rec, httptest.NewRequest(method, "/page", nil)) require.Equal(t, http.StatusOK, rec.Code, "%s /page", method) } // The point of registering the pair rather than reaching for chi's // middleware.GetHead: the handler sees the method that arrived. assert.Equal(t, []string{http.MethodGet, http.MethodHead}, seen) } func TestGetHeadAnswersAHeadWithTheHeadersOfItsGetAndNoBody(t *testing.T) { r := chi.NewRouter() GetHead(r, "/page", func(w http.ResponseWriter, _ *http.Request) { w.Header().Set("Content-Type", "text/plain; charset=utf-8") _, _ = io.WriteString(w, "a body of some length") }) // A real server, because dropping the body of a HEAD is net/http's job and // a recorder does not do it. srv := httptest.NewServer(r) t.Cleanup(srv.Close) get, err := http.Get(srv.URL + "/page") require.NoError(t, err) t.Cleanup(func() { _ = get.Body.Close() }) getBody, err := io.ReadAll(get.Body) require.NoError(t, err) head, err := http.Head(srv.URL + "/page") require.NoError(t, err) t.Cleanup(func() { _ = head.Body.Close() }) headBody, err := io.ReadAll(head.Body) require.NoError(t, err) assert.Equal(t, get.StatusCode, head.StatusCode) assert.Equal(t, get.Header.Get("Content-Type"), head.Header.Get("Content-Type")) assert.Empty(t, headBody) // The length a HEAD promises is the one a GET delivers, counted by net/http // off the writes it then discards. assert.Equal(t, strconv.Itoa(len(getBody)), head.Header.Get("Content-Length")) } func TestGetHeadRegistersHeadInTheRoutingTree(t *testing.T) { r := chi.NewRouter() GetHead(r, "/page", func(http.ResponseWriter, *http.Request) {}) r.Post("/page", func(http.ResponseWriter, *http.Request) {}) // This is the question chi's middleware.GetHead leaves the tree unable to // answer, and the one the donors' every-GET-has-a-HEAD-twin test asks. methods := map[string]bool{} require.NoError(t, chi.Walk(r, func(method, route string, _ http.Handler, _ ...func(http.Handler) http.Handler) error { if route == "/page" { methods[method] = true } return nil })) assert.True(t, methods[http.MethodGet], "GET is registered") assert.True(t, methods[http.MethodHead], "HEAD is registered") } func TestGetHeadLeavesEveryOtherMethodUnregistered(t *testing.T) { r := chi.NewRouter() GetHead(r, "/page", func(http.ResponseWriter, *http.Request) {}) rec := httptest.NewRecorder() r.ServeHTTP(rec, httptest.NewRequest(http.MethodPost, "/page", nil)) assert.Equal(t, http.StatusMethodNotAllowed, rec.Code) } func TestRenderRefusalsAnswersAnUnroutedPathThroughTheRenderer(t *testing.T) { render, calls := recordRefusals() r := chi.NewRouter() RenderRefusals(r, render) GetHead(r, "/page", func(http.ResponseWriter, *http.Request) {}) rec := httptest.NewRecorder() r.ServeHTTP(rec, httptest.NewRequest(http.MethodGet, "/nowhere", nil)) assert.Equal(t, http.StatusNotFound, rec.Code) assert.Equal(t, pages.NotFoundMessage, rec.Body.String()) require.Len(t, *calls, 1) assert.Equal(t, refusal{ status: http.StatusNotFound, message: pages.NotFoundMessage, method: http.MethodGet, }, (*calls)[0]) } func TestRenderRefusalsAnswersAnUnallowedMethodThroughTheRenderer(t *testing.T) { render, calls := recordRefusals() r := chi.NewRouter() RenderRefusals(r, render) GetHead(r, "/page", func(http.ResponseWriter, *http.Request) {}) rec := httptest.NewRecorder() r.ServeHTTP(rec, httptest.NewRequest(http.MethodPost, "/page", nil)) assert.Equal(t, http.StatusMethodNotAllowed, rec.Code) assert.Equal(t, pages.MethodMessage, rec.Body.String()) require.Len(t, *calls, 1) assert.Equal(t, refusal{ status: http.StatusMethodNotAllowed, message: pages.MethodMessage, method: http.MethodPost, }, (*calls)[0]) } func TestRenderRefusalsIsInheritedBySubRouters(t *testing.T) { render, calls := recordRefusals() r := chi.NewRouter() RenderRefusals(r, render) r.Route("/repo", func(r chi.Router) { GetHead(r, "/", func(http.ResponseWriter, *http.Request) {}) }) rec := httptest.NewRecorder() r.ServeHTTP(rec, httptest.NewRequest(http.MethodGet, "/repo/nowhere", nil)) assert.Equal(t, http.StatusNotFound, rec.Code) require.Len(t, *calls, 1) } func TestRenderRefusalsRequiresARenderCallback(t *testing.T) { assert.PanicsWithValue(t, "chimw: RenderRefusals needs a render callback", func() { RenderRefusals(chi.NewRouter(), nil) }) } // TestRenderRefusalsTakesAMethodValueOfTheDonorsShape pins the seam down: the // services pass their existing renderError by name, with no closure per call // site, and that only keeps working while ErrorRenderer has that signature. func TestRenderRefusalsTakesAMethodValueOfTheDonorsShape(t *testing.T) { var s server r := chi.NewRouter() RenderRefusals(r, s.renderError) rec := httptest.NewRecorder() r.ServeHTTP(rec, httptest.NewRequest(http.MethodGet, "/nowhere", nil)) assert.Equal(t, http.StatusNotFound, rec.Code) assert.True(t, strings.HasPrefix(rec.Body.String(), "page:"), "body %q", rec.Body.String()) } type server struct{} func (s *server) renderError(w http.ResponseWriter, _ *http.Request, status int, message string) { w.WriteHeader(status) _, _ = io.WriteString(w, "page: "+message) }