diff --git a/group.go b/group.go index 28ff93f84..a06d92e2a 100644 --- a/group.go +++ b/group.go @@ -12,9 +12,10 @@ import ( // routes that share a common middleware or functionality that should be separate // from the parent echo instance while still inheriting from it. type Group struct { - echo *Echo - prefix string - middleware []MiddlewareFunc + echo *Echo + prefix string + middleware []MiddlewareFunc + notFoundRoutes map[string]groupNotFoundRoute // noAutoRegisterRoutes is a flag that indicates whether Group should NOT register 404 routes automatically // when there are middlewares registered with the group. @@ -23,11 +24,17 @@ type Group struct { noAutoRegisterRoutes bool } +type groupNotFoundRoute struct { + route Route + path string +} + // Use implements `Echo#Use()` for sub-routes within the Group. // // Important! Group middlewares are executed in case there was no exact route match as by default Group registers // `/*` NotFound routes for itself. If this kind of behavior is not needed, then create an Echo instance with the ` noAutoRegisterRoutes ` // flag set to true. Example `echo.NewWithConfig(echo.Config{NoGroupAutoRegister404Routes: true})`. +// Explicit catch-all RouteNotFound handlers and their route-level middleware are preserved when Use is called again. func (g *Group) Use(middleware ...MiddlewareFunc) { g.middleware = append(g.middleware, middleware...) if len(g.middleware) == 0 { @@ -41,11 +48,19 @@ func (g *Group) Use(middleware ...MiddlewareFunc) { // So we register catch all route (404 is a safe way to emulate route match) for this group and now during routing the // Router would find route to match our request path and therefore guarantee the middleware(s) will get executed. // Note: we use nil handler so Router would choose the default 404 handler. This may not work with custom routers. - if _, err := g.AddRoute(Route{Method: RouteNotFound, Path: "", allowOverwrite: true}); err != nil { - panic(err) // this is how `v4` handles errors. `v5` has methods to have panic-free usage - } - if _, err := g.AddRoute(Route{Method: RouteNotFound, Path: "/*", allowOverwrite: true}); err != nil { - panic(err) // this is how `v4` handles errors. `v5` has methods to have panic-free usage + for _, path := range []string{"", "/*"} { + route := Route{Method: RouteNotFound, Path: path, allowOverwrite: true} + if existing, ok := g.notFoundRoutes[path]; ok { + if _, err := g.echo.Router().Routes().FindByMethodPath(RouteNotFound, existing.path); err == nil { + route = existing.route + route.allowOverwrite = true + } else { + delete(g.notFoundRoutes, path) + } + } + if _, err := g.AddRoute(route); err != nil { + panic(err) // this is how `v4` handles errors. `v5` has methods to have panic-free usage + } } } @@ -208,5 +223,13 @@ func (g *Group) AddRoute(route Route) (RouteInfo, error) { // multiple routes, which would lead to later add() calls overwriting the // middleware from earlier calls. groupRoute := route.WithPrefix(g.prefix, append([]MiddlewareFunc{}, g.middleware...)) - return g.echo.add(groupRoute) + ri, err := g.echo.add(groupRoute) + if err == nil && route.Method == RouteNotFound && route.Handler != nil && (route.Path == "" || route.Path == "/*") { + if g.notFoundRoutes == nil { + g.notFoundRoutes = make(map[string]groupNotFoundRoute) + } + route.Middlewares = append([]MiddlewareFunc(nil), route.Middlewares...) + g.notFoundRoutes[route.Path] = groupNotFoundRoute{route: route, path: ri.Path} + } + return ri, err } diff --git a/group_notfound_test.go b/group_notfound_test.go new file mode 100644 index 000000000..015a626dd --- /dev/null +++ b/group_notfound_test.go @@ -0,0 +1,199 @@ +// SPDX-License-Identifier: MIT +// SPDX-FileCopyrightText: © 2026 LabStack LLC and Echo contributors + +package echo + +import ( + "errors" + "io" + "net/http" + "net/http/httptest" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func groupNotFoundTraceMiddleware(name string) MiddlewareFunc { + return func(next HandlerFunc) HandlerFunc { + return func(c *Context) error { + c.Response().Header().Add("X-Trace", name) + return next(c) + } + } +} + +func TestGroupUsePreservesNotFoundHandler(t *testing.T) { + for _, path := range []string{"", "/*"} { + for _, order := range []string{"before middleware", "between middleware", "after middleware"} { + t.Run(path+"/"+order, func(t *testing.T) { + e := New() + g := e.Group("/api") + register := func() { + _, err := g.AddRoute(Route{ + Method: RouteNotFound, + Path: path, + Name: "custom-not-found", + Handler: func(c *Context) error { + c.Response().Header().Add("X-Trace", "handler") + return c.String(http.StatusNotFound, "custom group 404") + }, + Middlewares: []MiddlewareFunc{groupNotFoundTraceMiddleware("route")}, + }) + require.NoError(t, err) + } + if order == "before middleware" { + register() + } + g.Use(groupNotFoundTraceMiddleware("first")) + if order == "between middleware" { + register() + } + g.Use(groupNotFoundTraceMiddleware("second")) + if order == "after middleware" { + register() + } + url := "/api" + if path == "/*" { + url += "/missing" + } + rec := httptest.NewRecorder() + e.ServeHTTP(rec, httptest.NewRequest(http.MethodGet, url, nil)) + assert.Equal(t, http.StatusNotFound, rec.Code) + assert.Equal(t, "custom group 404", rec.Body.String()) + assert.Equal(t, []string{"first", "second", "route", "handler"}, rec.Header().Values("X-Trace")) + ri, err := e.Router().Routes().FindByMethodPath(RouteNotFound, "/api"+path) + require.NoError(t, err) + assert.Equal(t, "custom-not-found", ri.Name) + }) + } + } +} + +func TestGroupUsePreservesNotFoundWithoutRouterOverwrite(t *testing.T) { + e := NewWithConfig(Config{Router: NewRouter(RouterConfig{AllowOverwritingRoute: false})}) + g := e.Group("/api") + g.RouteNotFound("/*", func(c *Context) error { + return c.String(http.StatusNotFound, "original") + }) + _, err := g.AddRoute(Route{ + Method: RouteNotFound, + Path: "/*", + Handler: func(c *Context) error { return c.String(http.StatusNotFound, "rejected") }, + }) + require.Error(t, err) + g.Use(groupNotFoundTraceMiddleware("first")) + g.Use(groupNotFoundTraceMiddleware("second")) + rec := httptest.NewRecorder() + e.ServeHTTP(rec, httptest.NewRequest(http.MethodGet, "/api/missing", nil)) + assert.Equal(t, http.StatusNotFound, rec.Code) + assert.Equal(t, "original", rec.Body.String()) + assert.Equal(t, []string{"first", "second"}, rec.Header().Values("X-Trace")) +} + +func TestGroupUsePreservesNotFoundMiddlewareSnapshot(t *testing.T) { + e := New() + g := e.Group("/api") + middlewares := []MiddlewareFunc{groupNotFoundTraceMiddleware("original")} + _, err := g.AddRoute(Route{ + Method: RouteNotFound, + Path: "/*", + Handler: func(c *Context) error { return c.String(http.StatusNotFound, "custom") }, + Middlewares: middlewares, + }) + require.NoError(t, err) + middlewares[0] = groupNotFoundTraceMiddleware("mutated") + g.Use(groupNotFoundTraceMiddleware("group")) + rec := httptest.NewRecorder() + e.ServeHTTP(rec, httptest.NewRequest(http.MethodGet, "/api/missing", nil)) + assert.Equal(t, "custom", rec.Body.String()) + assert.Equal(t, []string{"group", "original"}, rec.Header().Values("X-Trace")) +} + +func TestGroupUseRejectsNotFoundRegistrationWithoutReplacingHandler(t *testing.T) { + e := New() + g := e.Group("/api") + g.RouteNotFound("/*", func(c *Context) error { return c.String(http.StatusNotFound, "original") }) + rejected := errors.New("registration rejected") + e.OnAddRoute = func(route Route) error { + if route.Name == "rejected" { + return rejected + } + return nil + } + _, err := g.AddRoute(Route{ + Method: RouteNotFound, + Path: "/*", + Name: "rejected", + Handler: func(c *Context) error { return c.String(http.StatusNotFound, "rejected") }, + }) + require.ErrorIs(t, err, rejected) + g.Use(groupNotFoundTraceMiddleware("group")) + rec := httptest.NewRecorder() + e.ServeHTTP(rec, httptest.NewRequest(http.MethodGet, "/api/missing", nil)) + assert.Equal(t, "original", rec.Body.String()) + assert.Equal(t, []string{"group"}, rec.Header().Values("X-Trace")) +} + +func TestGroupUseNotFoundWithAutoRegistrationDisabled(t *testing.T) { + e := NewWithConfig(Config{NoGroupAutoRegister404Routes: true}) + g := e.Group("/api") + g.RouteNotFound("/*", func(c *Context) error { return c.String(http.StatusNotFound, "custom") }, + groupNotFoundTraceMiddleware("route")) + g.Use(groupNotFoundTraceMiddleware("group")) + rec := httptest.NewRecorder() + e.ServeHTTP(rec, httptest.NewRequest(http.MethodGet, "/api/missing", nil)) + assert.Equal(t, "custom", rec.Body.String()) + assert.Equal(t, []string{"route"}, rec.Header().Values("X-Trace")) + _, err := e.Router().Routes().FindByMethodPath(RouteNotFound, "/api") + assert.Error(t, err) +} + +func TestGroupUseDoesNotRestoreRemovedNotFoundHandler(t *testing.T) { + e := NewWithConfig(Config{Router: NewRouter(RouterConfig{ + NotFoundHandler: func(c *Context) error { return c.String(http.StatusNotFound, "default") }, + })}) + g := e.Group("/api") + g.RouteNotFound("/*", func(c *Context) error { return c.String(http.StatusNotFound, "removed") }) + g.GET("/*", func(c *Context) error { return c.NoContent(http.StatusOK) }) + require.NoError(t, e.Router().Remove(RouteNotFound, "/api/*")) + require.NoError(t, e.Router().Remove(http.MethodGet, "/api/*")) + g.Use(groupNotFoundTraceMiddleware("first")) + g.Use(groupNotFoundTraceMiddleware("second")) + rec := httptest.NewRecorder() + e.ServeHTTP(rec, httptest.NewRequest(http.MethodGet, "/api/missing", nil)) + assert.Equal(t, "default", rec.Body.String()) + assert.Equal(t, []string{"first", "second"}, rec.Header().Values("X-Trace")) +} + +func TestGroupUseNotFoundOverHTTP(t *testing.T) { + e := NewWithConfig(Config{Router: NewRouter(RouterConfig{ + NotFoundHandler: func(c *Context) error { return c.String(http.StatusNotFound, "default") }, + })}) + g := e.Group("/api") + g.RouteNotFound("/*", func(c *Context) error { return c.String(http.StatusNotFound, "custom") }) + g.Use(groupNotFoundTraceMiddleware("first")) + g.Use(groupNotFoundTraceMiddleware("second")) + server := httptest.NewServer(e) + t.Cleanup(server.Close) + client := server.Client() + t.Cleanup(client.CloseIdleConnections) + for _, tc := range []struct { + path string + body string + }{ + {path: "/api/missing", body: "custom"}, + {path: "/api", body: "default"}, + } { + t.Run(tc.path, func(t *testing.T) { + response, err := client.Get(server.URL + tc.path) + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, response.Body.Close()) }) + body, err := io.ReadAll(response.Body) + require.NoError(t, err) + assert.Equal(t, http.StatusNotFound, response.StatusCode) + assert.Equal(t, tc.body, string(body)) + assert.Equal(t, []string{"first", "second"}, response.Header.Values("X-Trace")) + }) + } +}