diff --git a/router.go b/router.go index 99950aae7..c30c8aa4f 100644 --- a/router.go +++ b/router.go @@ -536,6 +536,9 @@ func (r *DefaultRouter) Add(route Route) (RouteInfo, error) { } paramNames := make([]string, 0) + // Positions of parameter markers after names are removed. Literal colons + // remain ordinary path bytes, so no sentinel byte is reserved. + paramMarkers := make([]int, 0) originalPath := path wasAdded := false var ri RouteInfo @@ -549,11 +552,12 @@ func (r *DefaultRouter) Add(route Route) (RouteInfo, error) { } j := i + 1 - r.insert(staticKind, path[:i], method, routeMethod{RouteInfo: &RouteInfo{Method: method}}) + r.insert(staticKind, path[:i], method, routeMethod{RouteInfo: &RouteInfo{Method: method}}, paramMarkers) for ; i < lcpIndex && path[i] != '/'; i++ { } paramNames = append(paramNames, path[j:i]) + paramMarkers = append(paramMarkers, j-1) path = path[:j] + path[i:] i, lcpIndex = j, len(path) @@ -566,14 +570,14 @@ func (r *DefaultRouter) Add(route Route) (RouteInfo, error) { orgRouteInfo: ri, wrappedHeadHandler: headH, } - r.insert(paramKind, path[:i], method, rm) + r.insert(paramKind, path[:i], method, rm, paramMarkers) wasAdded = true break } else { - r.insert(paramKind, path[:i], method, routeMethod{RouteInfo: &RouteInfo{Method: method}}) + r.insert(paramKind, path[:i], method, routeMethod{RouteInfo: &RouteInfo{Method: method}}, paramMarkers) } } else if path[i] == anyLabel { - r.insert(staticKind, path[:i], method, routeMethod{RouteInfo: &RouteInfo{Method: method}}) + r.insert(staticKind, path[:i], method, routeMethod{RouteInfo: &RouteInfo{Method: method}}, paramMarkers) paramNames = append(paramNames, "*") ri = route.ToRouteInfo(paramNames) rm := routeMethod{ @@ -582,7 +586,7 @@ func (r *DefaultRouter) Add(route Route) (RouteInfo, error) { orgRouteInfo: ri, wrappedHeadHandler: headH, } - r.insert(anyKind, path[:i+1], method, rm) + r.insert(anyKind, path[:i+1], method, rm, paramMarkers) wasAdded = true break } @@ -596,7 +600,7 @@ func (r *DefaultRouter) Add(route Route) (RouteInfo, error) { orgRouteInfo: ri, wrappedHeadHandler: headH, } - r.insert(staticKind, path, method, rm) + r.insert(staticKind, path, method, rm, paramMarkers) } r.storeRouteInfo(ri) @@ -623,12 +627,13 @@ func (r *DefaultRouter) storeRouteInfo(ri RouteInfo) { r.routes = append(r.routes, ri) } -func (r *DefaultRouter) insert(t kind, path string, method string, ri routeMethod) { +func (r *DefaultRouter) insert(t kind, path string, method string, ri routeMethod, paramMarkers []int) { if len(ri.Parameters) > r.maxPathParamsLength { r.maxPathParamsLength = len(ri.Parameters) } currentNode := r.tree // Current node as root search := path + searchOffset := 0 for { searchLen := len(search) @@ -717,8 +722,16 @@ func (r *DefaultRouter) insert(t kind, path string, method string, ri routeMetho } currentNode.refreshLeaf() } else if lcpLen < searchLen { + searchOffset += lcpLen search = search[lcpLen:] - c := currentNode.findChildWithLabel(search[0]) + isParamMarker := false + for _, marker := range paramMarkers { + if marker == searchOffset { + isParamMarker = true + break + } + } + c := currentNode.findChildWithLabel(search[0], isParamMarker) if c != nil { // Go deeper currentNode = c @@ -810,13 +823,13 @@ func (n *node) findStaticChild(l byte) *node { return nil } -func (n *node) findChildWithLabel(l byte) *node { +func (n *node) findChildWithLabel(l byte, isParamMarker bool) *node { + if isParamMarker { + return n.paramChild + } if c := n.findStaticChild(l); c != nil { return c } - if l == paramLabel { - return n.paramChild - } if l == anyLabel { return n.anyChild } @@ -937,8 +950,8 @@ func (r *DefaultRouter) Route(c *Context) HandlerFunc { searchIndex -= len(previous.prefix) } else { paramIndex-- - // for param/any node.prefix value is always `:` so we can not deduce searchIndex from that and must use pValue - // for that index as it would also contain part of path we cut off before moving into node we are backtracking from + // param/any node prefixes are a single marker byte, so restore searchIndex + // from the value stored for that param instead searchIndex -= len(pathValues[paramIndex].Value) pathValues[paramIndex].Value = "" } diff --git a/router_test.go b/router_test.go index 22c1dc759..fd7786f95 100644 --- a/router_test.go +++ b/router_test.go @@ -1320,6 +1320,106 @@ func TestRouterParamStaticConflict(t *testing.T) { } } +// Issue #3111 +func TestRouterParam_escapeColonAndParamConflict(t *testing.T) { + var testCases = []struct { + name string + routes []string + whenURL string + expectRoute string + expectParam map[string]string + }{ + { + name: "escaped colon route first, request escaped colon route", + routes: []string{`/name\:verb/x`, `/name:id`}, + whenURL: "/name:verb/x", + expectRoute: `/name\:verb/x`, + expectParam: map[string]string{}, + }, + { + name: "escaped colon route first, request param route", + routes: []string{`/name\:verb/x`, `/name:id`}, + whenURL: "/name1", + expectRoute: "/name:id", + expectParam: map[string]string{"id": "1"}, + }, + { + name: "param route first, request escaped colon route", + routes: []string{`/name:id`, `/name\:verb/x`}, + whenURL: "/name:verb/x", + expectRoute: `/name\:verb/x`, + expectParam: map[string]string{}, + }, + { + name: "param route first, request param route", + routes: []string{`/name:id`, `/name\:verb/x`}, + whenURL: "/name1", + expectRoute: "/name:id", + expectParam: map[string]string{"id": "1"}, + }, + } + for _, tc := range testCases { + t.Run(tc.name, func(t *testing.T) { + e := New() + for _, route := range tc.routes { + e.GET(route, handlerFunc) + } + + c := e.NewContext(httptest.NewRequest(http.MethodGet, tc.whenURL, nil), nil) + + handler := e.router.Route(c) + + assert.NoError(t, handler(c)) + assert.Equal(t, tc.expectRoute, c.Path()) + for param, expectedValue := range tc.expectParam { + assert.Equal(t, expectedValue, c.pathValues.GetOr(param, "---none---")) + } + checkUnusedParamValues(t, c, tc.expectParam) + }) + } +} + +func TestRouterParamLiteralByteConflictServeHTTP(t *testing.T) { + tests := []struct { + name, literalRoute, literalRequest string + }{ + {"escaped colon", `/name\:verb/x`, "/name:verb/x"}, + {"encoded NUL", "/name\x00verb/x", "/name%00verb/x"}, + } + for _, tc := range tests { + for _, literalFirst := range []bool{true, false} { + name := tc.name + "/parameter-first" + routes := []string{"/name:id", tc.literalRoute} + if literalFirst { + name = tc.name + "/literal-first" + routes[0], routes[1] = routes[1], routes[0] + } + t.Run(name, func(t *testing.T) { + e := New() + for _, route := range routes { + e.GET(route, func(c *Context) error { + return c.String(http.StatusOK, c.RouteInfo().Path) + }) + } + for _, request := range []struct{ path, want string }{ + {tc.literalRequest, tc.literalRoute}, + {"/name1", "/name:id"}, + } { + t.Run(request.path, func(t *testing.T) { + rec := httptest.NewRecorder() + req := httptest.NewRequest(http.MethodGet, request.path, nil) + if !assert.NotPanics(t, func() { e.ServeHTTP(rec, req) }) { + return + } + assert.Equal(t, http.StatusOK, rec.Code) + assert.Equal(t, request.want, rec.Body.String()) + }) + } + }) + } + } +} + func TestRouterParam_escapeColon(t *testing.T) { // to allow Google cloud API like route paths with colon in them // i.e. https://service.name/v1/some/resource/name:customVerb <- that `:customVerb` is not path param. It is just a string