package correlation import ( "net/http" "net/http/httptest" "testing" ) func TestMiddlewarePreservesValidCorrelationID(t *testing.T) { handler := Middleware(http.HandlerFunc(func(response http.ResponseWriter, request *http.Request) { if got := FromContext(request.Context()); got != "request-123" { t.Errorf("context correlation ID = %q", got) } response.WriteHeader(http.StatusNoContent) })) request := httptest.NewRequest(http.MethodGet, "/", nil) request.Header.Set(Header, "request-123") response := httptest.NewRecorder() handler.ServeHTTP(response, request) if response.Header().Get(Header) != "request-123" { t.Fatalf("response correlation ID = %q", response.Header().Get(Header)) } } func TestMiddlewareReplacesInvalidCorrelationID(t *testing.T) { handler := Middleware(http.HandlerFunc(func(response http.ResponseWriter, request *http.Request) { if len(FromContext(request.Context())) < 8 { t.Error("generated correlation ID is too short") } response.WriteHeader(http.StatusNoContent) })) request := httptest.NewRequest(http.MethodGet, "/", nil) request.Header.Set(Header, "secret\nforged") response := httptest.NewRecorder() handler.ServeHTTP(response, request) if response.Header().Get(Header) == "secret\nforged" || response.Header().Get(Header) == "" { t.Fatalf("invalid correlation ID was not replaced: %q", response.Header().Get(Header)) } }