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)
}