Skip to content
Merged
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
41 changes: 27 additions & 14 deletions router.go
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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)

Expand All @@ -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{
Expand All @@ -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
}
Expand All @@ -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)
Expand All @@ -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)
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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
}
Expand Down Expand Up @@ -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 = ""
}
Expand Down
100 changes: 100 additions & 0 deletions router_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
Loading