Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
11 changes: 10 additions & 1 deletion middleware/path_rewrite.go
Original file line number Diff line number Diff line change
Expand Up @@ -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)
})
}
Expand Down
89 changes: 89 additions & 0 deletions middleware/path_rewrite_absolute_probe_test.go
Original file line number Diff line number Diff line change
@@ -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)
}
})
}
}
75 changes: 75 additions & 0 deletions middleware/path_rewrite_mount_test.go
Original file line number Diff line number Diff line change
@@ -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)
}
})
}
}