diff --git a/internal/middleware/request_id_internal_test.go b/internal/middleware/request_id_internal_test.go new file mode 100644 index 0000000..5b381d9 --- /dev/null +++ b/internal/middleware/request_id_internal_test.go @@ -0,0 +1,118 @@ +package middleware + +import ( + "log/slog" + "net/http" + "net/http/httptest" + "os" + "strings" + "testing" + + "github.com/go-chi/chi/v5/middleware" + + "sneak.berlin/go/pixa/internal/config" +) + +// sendRequestID sends a request through the RequestID middleware, carrying +// incoming as its X-Request-Id unless that is empty. It returns the ID the +// next handler found in the request context, which the upstream fetch sends +// and the log lines carry, and the X-Request-Id of the response. +func sendRequestID(t *testing.T, incoming string) (string, string) { + t.Helper() + + mw := &Middleware{log: slog.Default(), config: &config.Config{}} + + var inContext string + + handler := mw.RequestID()(http.HandlerFunc( + func(_ http.ResponseWriter, r *http.Request) { + inContext = middleware.GetReqID(r.Context()) + })) + + req := httptest.NewRequestWithContext(t.Context(), http.MethodGet, "/", nil) + if incoming != "" { + req.Header.Set("X-Request-Id", incoming) + } + + rec := httptest.NewRecorder() + handler.ServeHTTP(rec, req) + + return inContext, rec.Header().Get("X-Request-Id") +} + +// TestRequestIDKeepsShortPlainID verifies that a request's own X-Request-Id of +// at most 64 letters, digits, '-', '_' or '.' is kept as its ID. +func TestRequestIDKeepsShortPlainID(t *testing.T) { + t.Parallel() + + for _, incoming := range []string{ + "client-request-id", + "A1_b2.c3-d4", + strings.Repeat("a", 64), + } { + inContext, inResponse := sendRequestID(t, incoming) + + if inContext != incoming || inResponse != incoming { + t.Errorf("incoming %q: context has %q, response %q, want both %q", + incoming, inContext, inResponse, incoming) + } + } +} + +// TestRequestIDReplacesLongOrUnusualID verifies that a request's own +// X-Request-Id that is over 64 characters or holds anything but letters, +// digits, '-', '_' or '.' is neither sent back nor sent upstream: the request +// gets a fresh ID instead. +func TestRequestIDReplacesLongOrUnusualID(t *testing.T) { + t.Parallel() + + for _, incoming := range []string{ + strings.Repeat("a", 65), + strings.Repeat("a", 9000), + "has space", + "a/b", + "a,b", + "