package chimw
import (
"bytes"
"context"
"encoding/json"
"io"
"log/slog"
"net/http"
"net/http/httptest"
"strings"
"testing"
"github.com/go-chi/chi/v5"
chimiddleware "github.com/go-chi/chi/v5/middleware"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"sourcecraft.dev/bigbes/sr-ht-ecore/middleware"
)
// capture builds a logger writing JSON records into a buffer, so a test can
// assert on the fields an operator would filter by rather than on a line of
// text.
func capture() (*slog.Logger, *bytes.Buffer) {
var buf bytes.Buffer
logger := slog.New(slog.NewJSONHandler(&buf, &slog.HandlerOptions{Level: slog.LevelDebug}))
return logger, &buf
}
// records parses everything written to the buffer.
func records(t *testing.T, buf *bytes.Buffer) []map[string]any {
t.Helper()
var out []map[string]any
for _, line := range strings.Split(strings.TrimSpace(buf.String()), "\n") {
if line == "" {
continue
}
var record map[string]any
require.NoError(t, json.Unmarshal([]byte(line), &record), "line %q", line)
out = append(out, record)
}
return out
}
// only asserts that exactly one record was written and returns it.
func only(t *testing.T, buf *bytes.Buffer) map[string]any {
t.Helper()
got := records(t, buf)
require.Len(t, got, 1, "one record per request")
return got[0]
}
// serve runs one request through a router carrying the logger and the given
// handler at /page.
func serve(f SlogFormatter, req *http.Request, h http.HandlerFunc) *httptest.ResponseRecorder {
r := chi.NewRouter()
r.Use(RequestLogger(f))
GetHead(r, "/page", h)
rec := httptest.NewRecorder()
r.ServeHTTP(rec, req)
return rec
}
func TestRequestLoggerWritesOneRecordWithTheRequestsFields(t *testing.T) {
logger, buf := capture()
rec := serve(SlogFormatter{Logger: logger},
httptest.NewRequest(http.MethodGet, "/page?q=secret", nil),
func(w http.ResponseWriter, _ *http.Request) {
_, _ = io.WriteString(w, "twelve bytes")
})
require.Equal(t, http.StatusOK, rec.Code)
record := only(t, buf)
assert.Equal(t, "request", record["msg"])
assert.Equal(t, "INFO", record["level"])
assert.Equal(t, http.MethodGet, record["method"])
assert.Equal(t, float64(http.StatusOK), record["status"])
assert.Equal(t, float64(len("twelve bytes")), record["bytes"])
assert.Greater(t, record["duration"], float64(0))
assert.NotContains(t, record, "request_id", "no RequestID middleware is installed")
// The path and not the request URI: the query string of these services
// carries what a viewer typed.
assert.Equal(t, "/page", record["path"])
}
func TestRequestLoggerLogsTheMethodThatArrived(t *testing.T) {
logger, buf := capture()
serve(SlogFormatter{Logger: logger},
httptest.NewRequest(http.MethodHead, "/page", nil),
func(w http.ResponseWriter, _ *http.Request) {
_, _ = io.WriteString(w, "body")
})
assert.Equal(t, http.MethodHead, only(t, buf)["method"])
}
func TestRequestLoggerPromotesAServerErrorAndNothingElse(t *testing.T) {
for _, tc := range []struct {
name string
status int
level string
}{
{"a page", http.StatusOK, "INFO"},
{"a redirect", http.StatusFound, "INFO"},
{"a refusal", http.StatusNotFound, "INFO"},
{"a forbidden page", http.StatusForbidden, "INFO"},
// 499 is below 500, which is the whole reason that number was picked:
// a client that hung up must not read as a server error.
{"a client that hung up", middleware.StatusClientClosedRequest, "INFO"},
{"a bug", http.StatusInternalServerError, "ERROR"},
{"a dependency that is down", http.StatusServiceUnavailable, "ERROR"},
} {
t.Run(tc.name, func(t *testing.T) {
logger, buf := capture()
serve(SlogFormatter{Logger: logger},
httptest.NewRequest(http.MethodGet, "/page", nil),
func(w http.ResponseWriter, _ *http.Request) { w.WriteHeader(tc.status) })
record := only(t, buf)
assert.Equal(t, float64(tc.status), record["status"])
assert.Equal(t, tc.level, record["level"])
})
}
}
func TestRequestLoggerReportsAHandlerThatWroteNothingAs200(t *testing.T) {
logger, buf := capture()
serve(SlogFormatter{Logger: logger},
httptest.NewRequest(http.MethodGet, "/page", nil),
func(http.ResponseWriter, *http.Request) {})
record := only(t, buf)
assert.Equal(t, float64(http.StatusOK), record["status"], "net/http answers an empty handler with 200")
assert.Equal(t, "INFO", record["level"])
}
func TestRequestLoggerReportsACancelledRequestAs499(t *testing.T) {
logger, buf := capture()
ctx, cancel := context.WithCancel(context.Background())
cancel()
req := httptest.NewRequest(http.MethodGet, "/page", nil).WithContext(ctx)
// The handler gave up without writing, because there was nobody left to
// write to.
serve(SlogFormatter{Logger: logger}, req, func(http.ResponseWriter, *http.Request) {})
record := only(t, buf)
assert.Equal(t, float64(middleware.StatusClientClosedRequest), record["status"])
assert.Equal(t, "INFO", record["level"], "a client that hung up must not page anybody")
}
func TestRequestLoggerCarriesTheRequestIDWhenChiSetsOne(t *testing.T) {
logger, buf := capture()
r := chi.NewRouter()
r.Use(chimiddleware.RequestID)
r.Use(RequestLogger(SlogFormatter{Logger: logger}))
GetHead(r, "/page", func(w http.ResponseWriter, r *http.Request) {
_, _ = io.WriteString(w, chimiddleware.GetReqID(r.Context()))
})
rec := httptest.NewRecorder()
r.ServeHTTP(rec, httptest.NewRequest(http.MethodGet, "/page", nil))
record := only(t, buf)
require.Contains(t, record, "request_id")
assert.Equal(t, rec.Body.String(), record["request_id"], "the id the handler saw")
assert.NotEmpty(t, record["request_id"])
}
func TestRequestLoggerSkipsWhatThePredicateSkips(t *testing.T) {
logger, buf := capture()
r := chi.NewRouter()
r.Use(RequestLogger(SlogFormatter{Logger: logger, Skip: SkipPaths("/healthz")}))
GetHead(r, "/healthz", func(w http.ResponseWriter, _ *http.Request) {
_, _ = io.WriteString(w, "ok")
})
GetHead(r, "/page", func(w http.ResponseWriter, _ *http.Request) {
_, _ = io.WriteString(w, "page")
})
// A query string does not turn a probe into a page.
r.ServeHTTP(httptest.NewRecorder(), httptest.NewRequest(http.MethodGet, "/healthz?probe=1", nil))
assert.Empty(t, records(t, buf))
r.ServeHTTP(httptest.NewRecorder(), httptest.NewRequest(http.MethodGet, "/page", nil))
assert.Equal(t, "/page", only(t, buf)["path"])
}
func TestRequestLoggerReportsAPanicOnASkippedPath(t *testing.T) {
logger, buf := capture()
r := chi.NewRouter()
r.Use(RequestLogger(SlogFormatter{Logger: logger, Skip: SkipPaths("/healthz")}))
r.Use(chimiddleware.Recoverer)
GetHead(r, "/healthz", func(http.ResponseWriter, *http.Request) {
panic("the probe is the bug")
})
rec := httptest.NewRecorder()
r.ServeHTTP(rec, httptest.NewRequest(http.MethodGet, "/healthz", nil))
require.Equal(t, http.StatusInternalServerError, rec.Code)
record := only(t, buf)
assert.Equal(t, "panic serving a request", record["msg"])
assert.Equal(t, "ERROR", record["level"])
assert.Equal(t, "the probe is the bug", record["panic"])
assert.Contains(t, record["stack"], "chimw")
assert.Equal(t, "/healthz", record["path"])
}
func TestRequestLoggerReportsAPanicAndStillWritesTheRequestLine(t *testing.T) {
logger, buf := capture()
r := chi.NewRouter()
r.Use(RequestLogger(SlogFormatter{Logger: logger}))
r.Use(chimiddleware.Recoverer)
GetHead(r, "/page", func(http.ResponseWriter, *http.Request) {
panic("boom")
})
rec := httptest.NewRecorder()
r.ServeHTTP(rec, httptest.NewRequest(http.MethodGet, "/page", nil))
require.Equal(t, http.StatusInternalServerError, rec.Code)
got := records(t, buf)
require.Len(t, got, 2)
assert.Equal(t, "panic serving a request", got[0]["msg"])
assert.Equal(t, "request", got[1]["msg"])
// The logger sits outside the recovery, so the line reports the status the
// viewer actually got.
assert.Equal(t, float64(http.StatusInternalServerError), got[1]["status"])
assert.Equal(t, "ERROR", got[1]["level"])
}
// TestRequestLoggerOutsideRecoverPanicsLogsTheStatusTheViewerGot pins the order
// the package doc asks for: the logger goes outermost, ahead of
// middleware.RecoverPanics, so the line reports the error page that was rendered
// rather than the nothing an unwinding stack has written so far.
func TestRequestLoggerOutsideRecoverPanicsLogsTheStatusTheViewerGot(t *testing.T) {
logger, buf := capture()
// RecoverPanics logs the panic and the stack through the default logger.
// Point that somewhere else so this test asserts on one buffer.
panics, _ := capture()
previous := slog.Default()
slog.SetDefault(panics)
t.Cleanup(func() { slog.SetDefault(previous) })
r := chi.NewRouter()
r.Use(RequestLogger(SlogFormatter{Logger: logger}))
r.Use(middleware.RecoverPanics(func(w http.ResponseWriter, _ *http.Request, _ any) {
w.WriteHeader(http.StatusInternalServerError)
_, _ = io.WriteString(w, "error page")
}))
GetHead(r, "/page", func(http.ResponseWriter, *http.Request) {
panic("boom")
})
rec := httptest.NewRecorder()
r.ServeHTTP(rec, httptest.NewRequest(http.MethodGet, "/page", nil))
require.Equal(t, http.StatusInternalServerError, rec.Code)
require.Equal(t, "error page", rec.Body.String())
record := only(t, buf)
assert.Equal(t, float64(http.StatusInternalServerError), record["status"])
assert.Equal(t, float64(len("error page")), record["bytes"])
assert.Equal(t, "ERROR", record["level"])
}
func TestRequestLoggerFallsBackToTheDefaultLogger(t *testing.T) {
logger, buf := capture()
previous := slog.Default()
slog.SetDefault(logger)
t.Cleanup(func() { slog.SetDefault(previous) })
// The zero value: no logger named, no message, nothing skipped.
serve(SlogFormatter{}, httptest.NewRequest(http.MethodGet, "/page", nil),
func(w http.ResponseWriter, _ *http.Request) { w.WriteHeader(http.StatusTeapot) })
record := only(t, buf)
assert.Equal(t, "request", record["msg"])
assert.Equal(t, float64(http.StatusTeapot), record["status"])
}
func TestRequestLoggerReadsTheDefaultLoggerAtWriteTime(t *testing.T) {
r := chi.NewRouter()
r.Use(RequestLogger(SlogFormatter{}))
GetHead(r, "/page", func(http.ResponseWriter, *http.Request) {})
// The router was built before the service installed its handler, which is
// the order a daemon does it in.
logger, buf := capture()
previous := slog.Default()
slog.SetDefault(logger)
t.Cleanup(func() { slog.SetDefault(previous) })
r.ServeHTTP(httptest.NewRecorder(), httptest.NewRequest(http.MethodGet, "/page", nil))
assert.Equal(t, "request", only(t, buf)["msg"])
}
func TestRequestLoggerTakesTheMessageItIsGiven(t *testing.T) {
logger, buf := capture()
serve(SlogFormatter{Logger: logger, Message: "http"},
httptest.NewRequest(http.MethodGet, "/page", nil),
func(http.ResponseWriter, *http.Request) {})
assert.Equal(t, "http", only(t, buf)["msg"])
}
func TestRequestLoggerLogsARefusalItNeverRouted(t *testing.T) {
logger, buf := capture()
render, _ := recordRefusals()
r := chi.NewRouter()
r.Use(RequestLogger(SlogFormatter{Logger: logger}))
RenderRefusals(r, render)
GetHead(r, "/page", func(http.ResponseWriter, *http.Request) {})
rec := httptest.NewRecorder()
r.ServeHTTP(rec, httptest.NewRequest(http.MethodGet, "/nowhere", nil))
require.Equal(t, http.StatusNotFound, rec.Code)
record := only(t, buf)
assert.Equal(t, float64(http.StatusNotFound), record["status"])
assert.Equal(t, "/nowhere", record["path"])
assert.Equal(t, "INFO", record["level"])
}
func TestSkipPathsMatchesNothingWhenGivenNothing(t *testing.T) {
skip := SkipPaths()
assert.False(t, skip(httptest.NewRequest(http.MethodGet, "/healthz", nil)))
}