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", "