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