diff --git a/bind.go b/bind.go index b0331a88a..13b6172a9 100644 --- a/bind.go +++ b/bind.go @@ -78,20 +78,16 @@ func BindBody(c *Context, target any) (err error) { switch mediatype { case MIMEApplicationJSON: if err = c.Echo().JSONSerializer.Deserialize(c, target); err != nil { - var hErr *HTTPError - if errors.As(err, &hErr) { - return err - } - return ErrBadRequest.Wrap(err) + return wrapBindBodyError(err) } case MIMEApplicationXML, MIMETextXML: if err = xml.NewDecoder(req.Body).Decode(target); err != nil { - return ErrBadRequest.Wrap(err) + return wrapBindBodyError(err) } case MIMEApplicationForm: params, err := c.FormValues() if err != nil { - return ErrBadRequest.Wrap(err) + return wrapBindBodyError(err) } if err = bindData(target, params, "form", nil); err != nil { return ErrBadRequest.Wrap(err) @@ -99,7 +95,7 @@ func BindBody(c *Context, target any) (err error) { case MIMEMultipartForm: params, err := c.MultipartForm() if err != nil { - return ErrBadRequest.Wrap(err) + return wrapBindBodyError(err) } if err = bindData(target, params.Value, "form", params.File); err != nil { return ErrBadRequest.Wrap(err) @@ -110,6 +106,13 @@ func BindBody(c *Context, target any) (err error) { return nil } +func wrapBindBodyError(err error) error { + if StatusCode(err) != 0 { + return err + } + return ErrBadRequest.Wrap(err) +} + // BindHeaders binds HTTP headers to a bindable object func BindHeaders(c *Context, target any) error { if err := bindData(target, c.Request().Header, "header", nil); err != nil { diff --git a/json.go b/json.go index d5fc9294c..eebe47852 100644 --- a/json.go +++ b/json.go @@ -54,6 +54,9 @@ func (d DefaultJSONSerializer) Deserialize(c *Context, target any) error { } }() if _, err := buf.ReadFrom(c.Request().Body); err != nil { + if StatusCode(err) != 0 { + return err + } return ErrBadRequest.Wrap(err) } if err := json.Unmarshal(buf.Bytes(), target); err != nil { diff --git a/middleware/body_limit.go b/middleware/body_limit.go index 4f1963e18..407f54f2e 100644 --- a/middleware/body_limit.go +++ b/middleware/body_limit.go @@ -4,9 +4,8 @@ package middleware import ( + "errors" "io" - "net/http" - "sync" "github.com/labstack/echo/v5" ) @@ -20,10 +19,14 @@ type BodyLimitConfig struct { LimitBytes int64 } +// limitedReader returns Echo's status-coded 413 error. Unlike +// http.MaxBytesReader, it does not tell net/http to close the connection +// after an over-limit read. type limitedReader struct { BodyLimitConfig reader io.ReadCloser read int64 + err error } // BodyLimit returns a BodyLimit middleware. @@ -45,14 +48,12 @@ func BodyLimitWithConfig(config BodyLimitConfig) echo.MiddlewareFunc { // ToMiddleware converts BodyLimitConfig to middleware or returns an error for invalid configuration func (config BodyLimitConfig) ToMiddleware() (echo.MiddlewareFunc, error) { + if config.LimitBytes < 0 { + return nil, errors.New("body limit must be non-negative") + } if config.Skipper == nil { config.Skipper = DefaultSkipper } - pool := sync.Pool{ - New: func() any { - return &limitedReader{BodyLimitConfig: config} - }, - } return func(next echo.HandlerFunc) echo.HandlerFunc { return func(c *echo.Context) error { @@ -66,14 +67,9 @@ func (config BodyLimitConfig) ToMiddleware() (echo.MiddlewareFunc, error) { return echo.ErrStatusRequestEntityTooLarge } - // Based on content read - r, ok := pool.Get().(*limitedReader) - if !ok { - return echo.NewHTTPError(http.StatusInternalServerError, "invalid pool object") - } - r.Reset(req.Body) - defer pool.Put(r) - req.Body = r + // Keep the wrapper attached to the request for its entire lifetime. + // Outer middleware may still use req.Body after next returns. + req.Body = &limitedReader{BodyLimitConfig: config, reader: req.Body} return next(c) } @@ -81,19 +77,37 @@ func (config BodyLimitConfig) ToMiddleware() (echo.MiddlewareFunc, error) { } func (r *limitedReader) Read(b []byte) (n int, err error) { + if r.err != nil { + return 0, r.err + } + if len(b) == 0 { + return 0, nil + } + remaining := r.LimitBytes - r.read + // If the caller asked for more bytes than are still allowed, cap the + // buffer one byte past the limit. That single extra byte is enough to + // tell whether the underlying reader holds more data than allowed, + // without ever reading more of it than necessary. + if int64(len(b))-1 > remaining { + b = b[:remaining+1] + } n, err = r.reader.Read(b) - r.read += int64(n) - if r.read > r.LimitBytes { - return n, echo.ErrStatusRequestEntityTooLarge + + if int64(n) <= remaining { + r.read += int64(n) + return n, err } - return + + // The underlying reader offered more data than the limit allows. Only + // hand out the allowed portion and make the error sticky, so callers + // that process the n>0 bytes before handling the error (as io.Reader + // documents) cannot read any further data on subsequent calls. + n = int(remaining) + r.read = r.LimitBytes + r.err = echo.ErrStatusRequestEntityTooLarge + return n, r.err } func (r *limitedReader) Close() error { return r.reader.Close() } - -func (r *limitedReader) Reset(reader io.ReadCloser) { - r.reader = reader - r.read = 0 -} diff --git a/middleware/body_limit_test.go b/middleware/body_limit_test.go index 68d904da8..06eb70cb6 100644 --- a/middleware/body_limit_test.go +++ b/middleware/body_limit_test.go @@ -6,6 +6,8 @@ package middleware import ( "bytes" "io" + "math" + "mime/multipart" "net/http" "net/http/httptest" "testing" @@ -41,8 +43,7 @@ func TestBodyLimitConfig_ToMiddleware(t *testing.T) { // Based on content read (overlimit) mw, err = BodyLimitConfig{LimitBytes: 2}.ToMiddleware() assert.NoError(t, err) - he := mw(h)(c).(echo.HTTPStatusCoder) - assert.Equal(t, http.StatusRequestEntityTooLarge, he.StatusCode()) + assert.Equal(t, http.StatusRequestEntityTooLarge, echo.StatusCode(mw(h)(c))) // Based on content read (within limit) req = httptest.NewRequest(http.MethodPost, "/", bytes.NewReader(hw)) @@ -64,8 +65,12 @@ func TestBodyLimitConfig_ToMiddleware(t *testing.T) { c = e.NewContext(req, rec) mw, err = BodyLimitConfig{LimitBytes: 2}.ToMiddleware() assert.NoError(t, err) - he = mw(h)(c).(echo.HTTPStatusCoder) - assert.Equal(t, http.StatusRequestEntityTooLarge, he.StatusCode()) + assert.Equal(t, http.StatusRequestEntityTooLarge, echo.StatusCode(mw(h)(c))) +} + +func TestBodyLimitRejectsNegativeLimit(t *testing.T) { + _, err := (BodyLimitConfig{LimitBytes: -1}).ToMiddleware() + assert.Error(t, err) } func TestBodyLimitAfterDecompressUsesDecodedSize(t *testing.T) { @@ -107,17 +112,200 @@ func TestBodyLimitReader(t *testing.T) { // read all should return ErrStatusRequestEntityTooLarge _, err := io.ReadAll(reader) - he := err.(echo.HTTPStatusCoder) - assert.Equal(t, http.StatusRequestEntityTooLarge, he.StatusCode()) + assert.ErrorIs(t, err, echo.ErrStatusRequestEntityTooLarge) - // reset reader and read two bytes must succeed + // A new request gets a new reader and can read within the limit. bt := make([]byte, 2) - reader.Reset(io.NopCloser(bytes.NewReader(hw))) + reader = &limitedReader{BodyLimitConfig: config, reader: io.NopCloser(bytes.NewReader(hw))} n, err := reader.Read(bt) assert.Equal(t, 2, n) assert.Equal(t, nil, err) } +func TestBodyLimitReader_singleOversizedRead(t *testing.T) { + hw := []byte("Hello, World!") + reader := &limitedReader{ + BodyLimitConfig: BodyLimitConfig{LimitBytes: 2}, + reader: io.NopCloser(bytes.NewReader(hw)), + } + + // a single read with a buffer much larger than the limit must deliver at + // most the allowed bytes and report the limit as exceeded + buf := make([]byte, 64) + n, err := reader.Read(buf) + assert.Equal(t, 2, n) + assert.ErrorIs(t, err, echo.ErrStatusRequestEntityTooLarge) +} + +func TestBodyLimitReader_noDataAfterLimitExceeded(t *testing.T) { + hw := bytes.Repeat([]byte("x"), 64) + reader := &limitedReader{ + BodyLimitConfig: BodyLimitConfig{LimitBytes: 5}, + reader: io.NopCloser(bytes.NewReader(hw)), + } + + // a caller following the io.Reader contract processes the n>0 bytes + // before considering the error and keeps calling Read; it must never + // receive more data once the limit has been exceeded + buf := make([]byte, 64) + total := 0 + var err error + for { + var n int + n, err = reader.Read(buf) + total += n + if n == 0 { + break + } + } + + assert.Equal(t, 5, total) + assert.ErrorIs(t, err, echo.ErrStatusRequestEntityTooLarge) +} + +func TestBodyLimitReader_exactLimitBody(t *testing.T) { + hw := []byte("ab") + reader := &limitedReader{ + BodyLimitConfig: BodyLimitConfig{LimitBytes: 2}, + reader: io.NopCloser(bytes.NewReader(hw)), + } + + // a body of exactly the limit size is not over the limit and must be + // readable completely + data, err := io.ReadAll(reader) + assert.NoError(t, err) + assert.Equal(t, "ab", string(data)) +} + +func TestBodyLimitReader_maxInt64Limit(t *testing.T) { + reader := &limitedReader{ + BodyLimitConfig: BodyLimitConfig{LimitBytes: math.MaxInt64}, + reader: io.NopCloser(bytes.NewReader([]byte("ok"))), + } + + data, err := io.ReadAll(reader) + assert.NoError(t, err) + assert.Equal(t, "ok", string(data)) +} + +type bodyLimitIntermittentReader struct{ calls int } + +func (r *bodyLimitIntermittentReader) Read(b []byte) (int, error) { + r.calls++ + if r.calls == 1 { + return 0, io.ErrUnexpectedEOF + } + return copy(b, "x"), nil +} + +func (r *bodyLimitIntermittentReader) Close() error { return nil } + +func TestBodyLimitReader_sourceErrorIsNotSticky(t *testing.T) { + reader := &limitedReader{ + BodyLimitConfig: BodyLimitConfig{LimitBytes: 5}, + reader: &bodyLimitIntermittentReader{}, + } + b := make([]byte, 1) + + n, err := reader.Read(b) + assert.Zero(t, n) + assert.ErrorIs(t, err, io.ErrUnexpectedEOF) + + n, err = reader.Read(b) + assert.NoError(t, err) + assert.Equal(t, 1, n) + assert.Equal(t, byte('x'), b[0]) +} + +func TestBodyLimitKeepsRequestReaderAfterHandler(t *testing.T) { + e := echo.New() + mw := BodyLimit(10) + readBody := func(body string) io.ReadCloser { + req := httptest.NewRequest(http.MethodPost, "/", bytes.NewBufferString(body)) + rec := httptest.NewRecorder() + c := e.NewContext(req, rec) + assert.NoError(t, mw(func(c *echo.Context) error { return nil })(c)) + return req.Body + } + + first := readBody("first") + second := readBody("second") + assert.NotSame(t, first, second) + data, err := io.ReadAll(first) + assert.NoError(t, err) + assert.Equal(t, "first", string(data)) +} + +func TestBodyLimit_oversizedBodyWithContractCompliantReader(t *testing.T) { + e := echo.New() + const limit = 5 + h := func(c *echo.Context) error { + buf := make([]byte, 64) + total := 0 + for { + n, err := c.Request().Body.Read(buf) + total += n + if err != nil { + assert.LessOrEqual(t, total, limit) + return err + } + if n == 0 { + break + } + } + return c.String(http.StatusOK, "ok") + } + mw, err := BodyLimitConfig{LimitBytes: limit}.ToMiddleware() + assert.NoError(t, err) + e.POST("/", h, mw) + + body := bytes.Repeat([]byte("x"), 10*limit) + req := httptest.NewRequest(http.MethodPost, "/", bytes.NewReader(body)) + req.ContentLength = -1 // force the content-read path + req.TransferEncoding = []string{"chunked"} + rec := httptest.NewRecorder() + e.ServeHTTP(rec, req) + assert.Equal(t, http.StatusRequestEntityTooLarge, rec.Code) +} + +func TestBodyLimitBindReturns413(t *testing.T) { + var multipartBody bytes.Buffer + multipartWriter := multipart.NewWriter(&multipartBody) + assert.NoError(t, multipartWriter.WriteField("x", "abcdefghij")) + assert.NoError(t, multipartWriter.Close()) + + tests := []struct { + name string + contentType string + body []byte + }{ + {"json", echo.MIMEApplicationJSON, []byte(`{"x":"abcdefghij"}`)}, + {"xml", echo.MIMEApplicationXML, []byte(`abcdefghij`)}, + {"form", echo.MIMEApplicationForm, []byte(`x=abcdefghij`)}, + {"multipart", multipartWriter.FormDataContentType(), multipartBody.Bytes()}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + e := echo.New() + e.POST("/", func(c *echo.Context) error { + var value struct { + X string `json:"x" xml:"x" form:"x"` + } + if err := c.Bind(&value); err != nil { + return err + } + return c.String(http.StatusOK, value.X) + }, BodyLimit(5)) + req := httptest.NewRequest(http.MethodPost, "/", bytes.NewReader(tt.body)) + req.Header.Set(echo.HeaderContentType, tt.contentType) + req.ContentLength = -1 + rec := httptest.NewRecorder() + e.ServeHTTP(rec, req) + assert.Equal(t, http.StatusRequestEntityTooLarge, rec.Code) + }) + } +} + func TestBodyLimit_skipper(t *testing.T) { e := echo.New() h := func(c *echo.Context) error {