diff --git a/go.mod b/go.mod new file mode 100644 index 0000000..f105ccf --- /dev/null +++ b/go.mod @@ -0,0 +1,17 @@ +module github.com/RoseMark45/echo + +go 1.23.0 + +require github.com/labstack/echo/v4 v4.13.4 + +require ( + github.com/labstack/gommon v0.4.2 // indirect + github.com/mattn/go-colorable v0.1.14 // indirect + github.com/mattn/go-isatty v0.0.20 // indirect + github.com/valyala/bytebufferpool v1.0.0 // indirect + github.com/valyala/fasttemplate v1.2.2 // indirect + golang.org/x/crypto v0.38.0 // indirect + golang.org/x/net v0.40.0 // indirect + golang.org/x/sys v0.33.0 // indirect + golang.org/x/text v0.25.0 // indirect +) diff --git a/go.sum b/go.sum new file mode 100644 index 0000000..b92f49c --- /dev/null +++ b/go.sum @@ -0,0 +1,29 @@ +github.com/davecgh/go-spew v1.1.1 h1:vj9j/u1bqnvCEfJOwUhtlOARqs3+rkHYY13jYWTU97c= +github.com/davecgh/go-spew v1.1.1/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38= +github.com/labstack/echo/v4 v4.13.4 h1:oTZZW+T3s9gAu5L8vmzihV7/lkXGZuITzTQkTEhcXEA= +github.com/labstack/echo/v4 v4.13.4/go.mod h1:g63b33BZ5vZzcIUF8AtRH40DrTlXnx4UMC8rBdndmjQ= +github.com/labstack/gommon v0.4.2 h1:F8qTUNXgG1+6WQmqoUWnz8WiEU60mXVVw0P4ht1WRA0= +github.com/labstack/gommon v0.4.2/go.mod h1:QlUFxVM+SNXhDL/Z7YhocGIBYOiwB0mXm1+1bAPHPyU= +github.com/mattn/go-colorable v0.1.14 h1:9A9LHSqF/7dyVVX6g0U9cwm9pG3kP9gSzcuIPHPsaIE= +github.com/mattn/go-colorable v0.1.14/go.mod h1:6LmQG8QLFO4G5z1gPvYEzlUgJ2wF+stgPZH1UqBm1s8= +github.com/mattn/go-isatty v0.0.20 h1:xfD0iDuEKnDkl03q4limB+vH+GxLEtL/jb4xVJSWWEY= +github.com/mattn/go-isatty v0.0.20/go.mod h1:W+V8PltTTMOvKvAeJH7IuucS94S2C6jfK/D7dTCTo3Y= +github.com/pmezard/go-difflib v1.0.0 h1:4DBwDE0NGyQoBHbLQYPwSUPoCMWR5BEzIk/f1lZbAQM= +github.com/pmezard/go-difflib v1.0.0/go.mod h1:iKH77koFhYxTK1pcRnkKkqfTogsbg7gZNVY4sRDYZ/4= +github.com/stretchr/testify v1.10.0 h1:Xv5erBjTwe/5IxqUQTdXv5kgmIvbHo3QQyRwhJsOfJA= +github.com/stretchr/testify v1.10.0/go.mod h1:r2ic/lqez/lEtzL7wO/rwa5dbSLXVDPFyf8C91i36aY= +github.com/valyala/bytebufferpool v1.0.0 h1:GqA5TC/0021Y/b9FG4Oi9Mr3q7XYx6KllzawFIhcdPw= +github.com/valyala/bytebufferpool v1.0.0/go.mod h1:6bBcMArwyJ5K/AmCkWv1jt77kVWyCJ6HpOuEn7z0Csc= +github.com/valyala/fasttemplate v1.2.2 h1:lxLXG0uE3Qnshl9QyaK6XJxMXlQZELvChBOCmQD0Loo= +github.com/valyala/fasttemplate v1.2.2/go.mod h1:KHLXt3tVN2HBp8eijSv/kGJopbvo7S+qRAEEKiv+SiQ= +golang.org/x/crypto v0.38.0 h1:jt+WWG8IZlBnVbomuhg2Mdq0+BBQaHbtqHEFEigjUV8= +golang.org/x/crypto v0.38.0/go.mod h1:MvrbAqul58NNYPKnOra203SB9vpuZW0e+RRZV+Ggqjw= +golang.org/x/net v0.40.0 h1:79Xs7wF06Gbdcg4kdCCIQArK11Z1hr5POQ6+fIYHNuY= +golang.org/x/net v0.40.0/go.mod h1:y0hY0exeL2Pku80/zKK7tpntoX23cqL3Oa6njdgRtds= +golang.org/x/sys v0.6.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= +golang.org/x/sys v0.33.0 h1:q3i8TbbEz+JRD9ywIRlyRAQbM0qF7hu24q3teo2hbuw= +golang.org/x/sys v0.33.0/go.mod h1:BJP2sWEmIv4KK5OTEluFJCKSidICx8ciO85XgH3Ak8k= +golang.org/x/text v0.25.0 h1:qVyWApTSYLk/drJRO5mDlNYskwQznZmkpV2c8q9zls4= +golang.org/x/text v0.25.0/go.mod h1:WEdwpYrmk1qmdHvhkSTNPm3app7v4rsT8F2UD6+VHIA= +gopkg.in/yaml.v3 v3.0.1 h1:fxVm/GzAzEWqLHuvctI91KS9hhNmmWOoWu0XTYJS7CA= +gopkg.in/yaml.v3 v3.0.1/go.mod h1:K4uyk7z7BCEPqu6E+C64Yfv1cQ7kz7rIZviUmN+EgEM= diff --git a/middleware/body_cache.go b/middleware/body_cache.go new file mode 100644 index 0000000..79813ac --- /dev/null +++ b/middleware/body_cache.go @@ -0,0 +1,95 @@ +package middleware + +import ( + "bytes" + "errors" + "io" + "net/http" + + "github.com/labstack/echo/v4" +) + +const DefaultBodyCacheKey = "rawBody" + +var ErrBodyTooLarge = errors.New("request body exceeds maximum cache size") + +type BodyCacheConfig struct { + // Limit is the maximum number of request-body bytes to cache. A zero value + // means there is no explicit limit. + Limit int64 + + // ContextKey is the Echo context key used to expose the cached raw body. + // When empty, DefaultBodyCacheKey is used. + ContextKey string +} + +func BodyCache() echo.MiddlewareFunc { + return BodyCacheWithConfig(BodyCacheConfig{}) +} + +func BodyCacheWithLimit(limit int64) echo.MiddlewareFunc { + return BodyCacheWithConfig(BodyCacheConfig{Limit: limit}) +} + +func BodyCacheWithConfig(config BodyCacheConfig) echo.MiddlewareFunc { + if config.ContextKey == "" { + config.ContextKey = DefaultBodyCacheKey + } + + return func(next echo.HandlerFunc) echo.HandlerFunc { + return func(c echo.Context) error { + req := c.Request() + if req == nil || req.Body == nil || req.Body == http.NoBody { + c.Set(config.ContextKey, []byte(nil)) + return next(c) + } + + body, err := readBodyWithLimit(req.Body, config.Limit) + if err != nil { + if errors.Is(err, ErrBodyTooLarge) { + return echo.NewHTTPError(http.StatusRequestEntityTooLarge, err.Error()) + } + return err + } + + c.Set(config.ContextKey, body) + req.Body = io.NopCloser(bytes.NewReader(body)) + + return next(c) + } + } +} + +func RestoreCachedBody(c echo.Context, key ...string) error { + contextKey := DefaultBodyCacheKey + if len(key) > 0 && key[0] != "" { + contextKey = key[0] + } + + body, ok := c.Get(contextKey).([]byte) + if !ok { + return errors.New("cached request body is not available") + } + + c.Request().Body = io.NopCloser(bytes.NewReader(body)) + return nil +} + +func readBodyWithLimit(body io.ReadCloser, limit int64) ([]byte, error) { + defer body.Close() + + if limit <= 0 { + return io.ReadAll(body) + } + + limited := io.LimitReader(body, limit+1) + readBody, err := io.ReadAll(limited) + if err != nil { + return nil, err + } + if int64(len(readBody)) > limit { + return nil, ErrBodyTooLarge + } + + return readBody, nil +} diff --git a/middleware/body_cache_test.go b/middleware/body_cache_test.go new file mode 100644 index 0000000..3498ba9 --- /dev/null +++ b/middleware/body_cache_test.go @@ -0,0 +1,168 @@ +package middleware + +import ( + "bytes" + "encoding/json" + "io" + "net/http" + "net/http/httptest" + "net/url" + "strings" + "testing" + + "github.com/labstack/echo/v4" +) + +type jsonPayload struct { + Foo string `json:"foo" xml:"foo" form:"foo"` +} + +func TestBodyCacheRestoresJSONAfterBind(t *testing.T) { + body := `{"foo":"bar"}` + + gotBody, status := bindThenReadBody(t, "application/json", body) + + if status != http.StatusOK { + t.Fatalf("expected status 200, got %d", status) + } + if gotBody != body { + t.Fatalf("expected downstream body %q, got %q", body, gotBody) + } +} + +func TestBodyCacheRestoresXMLAfterBind(t *testing.T) { + body := `bar` + + gotBody, status := bindThenReadBody(t, "application/xml", body) + + if status != http.StatusOK { + t.Fatalf("expected status 200, got %d", status) + } + if gotBody != body { + t.Fatalf("expected downstream body %q, got %q", body, gotBody) + } +} + +func TestBodyCacheRestoresFormAfterBind(t *testing.T) { + values := url.Values{"foo": []string{"bar"}} + body := values.Encode() + + gotBody, status := bindThenReadBody(t, echo.MIMEApplicationForm, body) + + if status != http.StatusOK { + t.Fatalf("expected status 200, got %d", status) + } + if gotBody != body { + t.Fatalf("expected downstream body %q, got %q", body, gotBody) + } +} + +func TestBodyCacheHandlesEmptyBody(t *testing.T) { + e := echo.New() + e.Use(BodyCache()) + e.POST("/", func(c echo.Context) error { + body, err := io.ReadAll(c.Request().Body) + if err != nil { + return err + } + return c.String(http.StatusOK, string(body)) + }) + + req := httptest.NewRequest(http.MethodPost, "/", nil) + rec := httptest.NewRecorder() + + e.ServeHTTP(rec, req) + + if rec.Code != http.StatusOK { + t.Fatalf("expected status 200, got %d", rec.Code) + } + if rec.Body.String() != "" { + t.Fatalf("expected empty body, got %q", rec.Body.String()) + } +} + +func TestBodyCacheRejectsBodyOverLimit(t *testing.T) { + e := echo.New() + e.Use(BodyCacheWithLimit(3)) + e.POST("/", func(c echo.Context) error { + return c.NoContent(http.StatusOK) + }) + + req := httptest.NewRequest(http.MethodPost, "/", strings.NewReader("abcd")) + rec := httptest.NewRecorder() + + e.ServeHTTP(rec, req) + + if rec.Code != http.StatusRequestEntityTooLarge { + t.Fatalf("expected status 413, got %d", rec.Code) + } +} + +func TestRestoreCachedBodyAllowsMultipleReads(t *testing.T) { + e := echo.New() + e.Use(BodyCache()) + e.POST("/", func(c echo.Context) error { + firstRead, err := io.ReadAll(c.Request().Body) + if err != nil { + return err + } + if err := RestoreCachedBody(c); err != nil { + return err + } + secondRead, err := io.ReadAll(c.Request().Body) + if err != nil { + return err + } + if !bytes.Equal(firstRead, secondRead) { + t.Fatalf("expected restored body %q, got %q", firstRead, secondRead) + } + return c.String(http.StatusOK, string(secondRead)) + }) + + req := httptest.NewRequest(http.MethodPost, "/", strings.NewReader(`{"foo":"bar"}`)) + rec := httptest.NewRecorder() + + e.ServeHTTP(rec, req) + + if rec.Code != http.StatusOK { + t.Fatalf("expected status 200, got %d", rec.Code) + } +} + +func bindThenReadBody(t *testing.T, contentType string, body string) (string, int) { + t.Helper() + + e := echo.New() + e.Use(BodyCache()) + e.Use(func(next echo.HandlerFunc) echo.HandlerFunc { + return func(c echo.Context) error { + var payload jsonPayload + if err := c.Bind(&payload); err != nil { + return err + } + if payload.Foo != "bar" { + encoded, _ := json.Marshal(payload) + return echo.NewHTTPError(http.StatusBadRequest, string(encoded)) + } + if err := RestoreCachedBody(c); err != nil { + return err + } + return next(c) + } + }) + e.POST("/", func(c echo.Context) error { + body, err := io.ReadAll(c.Request().Body) + if err != nil { + return err + } + return c.String(http.StatusOK, string(body)) + }) + + req := httptest.NewRequest(http.MethodPost, "/", strings.NewReader(body)) + req.Header.Set(echo.HeaderContentType, contentType) + rec := httptest.NewRecorder() + + e.ServeHTTP(rec, req) + + return rec.Body.String(), rec.Code +}