diff --git a/middleware/path_rewrite.go b/middleware/path_rewrite.go index 99af62c0c..dba5bd80c 100644 --- a/middleware/path_rewrite.go +++ b/middleware/path_rewrite.go @@ -3,13 +3,22 @@ package middleware import ( "net/http" "strings" + + "github.com/go-chi/chi/v5" ) -// PathRewrite is a simple middleware which allows you to rewrite the request URL path. +// PathRewrite replaces the first occurrence of old with new in the request URL path. +// When the routing path has already been set, it is rewritten independently. +// In a mounted router, the routing path is relative to the mount point, so old +// and new should be relative to that point to affect routing. To rewrite a path +// including the mount prefix, install PathRewrite on the parent router instead. func PathRewrite(old, new string) func(http.Handler) http.Handler { return func(next http.Handler) http.Handler { return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { r.URL.Path = strings.Replace(r.URL.Path, old, new, 1) + if rctx := chi.RouteContext(r.Context()); rctx != nil && rctx.RoutePath != "" { + rctx.RoutePath = strings.Replace(rctx.RoutePath, old, new, 1) + } next.ServeHTTP(w, r) }) } diff --git a/middleware/path_rewrite_absolute_probe_test.go b/middleware/path_rewrite_absolute_probe_test.go new file mode 100644 index 000000000..a1e57345b --- /dev/null +++ b/middleware/path_rewrite_absolute_probe_test.go @@ -0,0 +1,89 @@ +package middleware_test + +import ( + "fmt" + "io" + "net/http" + "net/http/httptest" + "testing" + + "github.com/go-chi/chi/v5" + "github.com/go-chi/chi/v5/middleware" +) + +// Full-path rewrites belong before the router that consumes the mount prefix. +func TestPathRewriteAbsoluteMountPlacement(t *testing.T) { + for _, placement := range []string{"parent", "mounted"} { + t.Run(placement, func(t *testing.T) { + root, sub := chi.NewRouter(), chi.NewRouter() + rewrite := middleware.PathRewrite("/api/old", "/api/new") + if placement == "parent" { + root.Use(rewrite) + } else { + sub.Use(rewrite) + } + sub.Get("/new/{id}", func(w http.ResponseWriter, r *http.Request) { + fmt.Fprintf(w, "%s|%s", r.URL.Path, chi.URLParam(r, "id")) + }) + sub.NotFound(func(w http.ResponseWriter, r *http.Request) { + w.WriteHeader(http.StatusNotFound) + fmt.Fprintf(w, "%s|%s", r.URL.Path, chi.RouteContext(r.Context()).RoutePath) + }) + root.Mount("/api", sub) + server := httptest.NewServer(root) + defer server.Close() + resp, err := server.Client().Get(server.URL + "/api/old/item") + if err != nil { + t.Fatal(err) + } + defer resp.Body.Close() + body, err := io.ReadAll(resp.Body) + if err != nil { + t.Fatal(err) + } + wantCode, wantBody := http.StatusOK, "/api/new/item|item" + if placement == "mounted" { + wantCode, wantBody = http.StatusNotFound, "/api/new/item|/old/item" + } + if resp.StatusCode != wantCode || string(body) != wantBody { + t.Fatalf("got %d %q, want %d %q", resp.StatusCode, body, wantCode, wantBody) + } + }) + } +} + +// URLFormat changes only the routing path. Do not restore its removed suffix +// by reconstructing RoutePath from URL.Path. +func TestPathRewriteAfterURLFormat(t *testing.T) { + for _, prefix := range []string{"", "/api"} { + t.Run(prefix, func(t *testing.T) { + sub := chi.NewRouter() + sub.Use(middleware.URLFormat) + sub.Use(middleware.PathRewrite("/old", "/new")) + sub.Get("/new/{id}", func(w http.ResponseWriter, r *http.Request) { + fmt.Fprintf(w, "%s|%s|%s|%s", r.URL.Path, chi.URLParam(r, "id"), chi.RouteContext(r.Context()).RoutePath, r.Context().Value(middleware.URLFormatCtxKey)) + }) + var handler http.Handler = sub + if prefix != "" { + root := chi.NewRouter() + root.Mount(prefix, sub) + handler = root + } + server := httptest.NewServer(handler) + defer server.Close() + resp, err := server.Client().Get(server.URL + prefix + "/old/item.json") + if err != nil { + t.Fatal(err) + } + defer resp.Body.Close() + body, err := io.ReadAll(resp.Body) + if err != nil { + t.Fatal(err) + } + want := prefix + "/new/item.json|item|/new/item|json" + if resp.StatusCode != http.StatusOK || string(body) != want { + t.Fatalf("got %d %q, want 200 %q", resp.StatusCode, body, want) + } + }) + } +} diff --git a/middleware/path_rewrite_mount_test.go b/middleware/path_rewrite_mount_test.go new file mode 100644 index 000000000..a407810b7 --- /dev/null +++ b/middleware/path_rewrite_mount_test.go @@ -0,0 +1,75 @@ +package middleware_test + +import ( + "fmt" + "io" + "net/http" + "net/http/httptest" + "testing" + + "github.com/go-chi/chi/v5" + "github.com/go-chi/chi/v5/middleware" +) + +func TestPathRewritePlainHandler(t *testing.T) { + handler := middleware.PathRewrite("/old", "/new")(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + fmt.Fprint(w, r.URL.Path) + })) + rec := httptest.NewRecorder() + handler.ServeHTTP(rec, httptest.NewRequest(http.MethodGet, "/old/item", nil)) + if rec.Code != http.StatusOK || rec.Body.String() != "/new/item" { + t.Fatalf("got %d %q", rec.Code, rec.Body.String()) + } +} + +func TestPathRewriteMountedRouter(t *testing.T) { + for _, tc := range []struct { + name string + prefix string + target string + want string + }{ + {"root", "", "/old/item?sort=asc", "/new/item|item|sort=asc"}, + {"mount", "/api", "/api/old/item?sort=asc", "/api/new/item|item|sort=asc"}, + {"nested mount", "/api/v1", "/api/v1/old/item?sort=asc", "/api/v1/new/item|item|sort=asc"}, + {"no replacement", "/api", "/api/new/item?sort=asc", "/api/new/item|item|sort=asc"}, + {"encoded parameter", "/api", "/api/old/a%2Fb?sort=asc", "/api/new/a/b|a%2Fb|sort=asc"}, + {"replace once", "/api", "/api/old/old?sort=asc", "/api/new/old|old|sort=asc"}, + {"mount prefix also matches", "/old", "/old/old/item?sort=asc", "/new/old/item|item|sort=asc"}, + } { + t.Run(tc.name, func(t *testing.T) { + sub := chi.NewRouter() + sub.Use(middleware.PathRewrite("/old", "/new")) + sub.Get("/new/{id}", func(w http.ResponseWriter, r *http.Request) { + fmt.Fprintf(w, "%s|%s|%s", r.URL.Path, chi.URLParam(r, "id"), r.URL.RawQuery) + if tc.prefix == "/old" && chi.RouteContext(r.Context()).RoutePath != "/new/item" { + t.Errorf("routing path = %q, want /new/item", chi.RouteContext(r.Context()).RoutePath) + } + }) + var handler http.Handler = sub + if tc.prefix != "" { + root := chi.NewRouter() + if tc.prefix == "/api/v1" { + root.Route("/api", func(r chi.Router) { r.Mount("/v1", sub) }) + } else { + root.Mount(tc.prefix, sub) + } + handler = root + } + server := httptest.NewServer(handler) + defer server.Close() + resp, err := server.Client().Get(server.URL + tc.target) + if err != nil { + t.Fatal(err) + } + defer resp.Body.Close() + body, err := io.ReadAll(resp.Body) + if err != nil { + t.Fatal(err) + } + if resp.StatusCode != http.StatusOK || string(body) != tc.want { + t.Fatalf("got %d %q, want 200 %q", resp.StatusCode, body, tc.want) + } + }) + } +}