diff --git a/bind.go b/bind.go index 1d4fe6f0a..00b57a5c0 100644 --- a/bind.go +++ b/bind.go @@ -95,6 +95,10 @@ func (b *DefaultBinder) BindBody(c Context, i interface{}) (err error) { } case MIMEApplicationXML, MIMETextXML: if err = xml.NewDecoder(req.Body).Decode(i); err != nil { + var httpError *HTTPError + if errors.As(err, &httpError) { + return httpError + } if ute, ok := err.(*xml.UnsupportedTypeError); ok { return NewHTTPError(http.StatusBadRequest, fmt.Sprintf("Unsupported type error: type=%v, error=%v", ute.Type, ute.Error())).SetInternal(err) } else if se, ok := err.(*xml.SyntaxError); ok { @@ -105,6 +109,10 @@ func (b *DefaultBinder) BindBody(c Context, i interface{}) (err error) { case MIMEApplicationForm: params, err := c.FormParams() if err != nil { + var httpError *HTTPError + if errors.As(err, &httpError) { + return httpError + } return NewHTTPError(http.StatusBadRequest, err.Error()).SetInternal(err) } if err = b.bindData(i, params, "form", nil); err != nil { @@ -113,6 +121,10 @@ func (b *DefaultBinder) BindBody(c Context, i interface{}) (err error) { case MIMEMultipartForm: params, err := c.MultipartForm() if err != nil { + var httpError *HTTPError + if errors.As(err, &httpError) { + return httpError + } return NewHTTPError(http.StatusBadRequest, err.Error()).SetInternal(err) } if err = b.bindData(i, params.Value, "form", params.File); err != nil { diff --git a/middleware/body_limit.go b/middleware/body_limit.go index e48aff61c..35a6a1826 100644 --- a/middleware/body_limit.go +++ b/middleware/body_limit.go @@ -6,8 +6,6 @@ package middleware import ( "fmt" "io" - "net/http" - "sync" "github.com/labstack/echo/v4" "github.com/labstack/gommon/bytes" @@ -28,6 +26,7 @@ type limitedReader struct { BodyLimitConfig reader io.ReadCloser read int64 + err error } // DefaultBodyLimitConfig is the default BodyLimit middleware config. @@ -58,11 +57,10 @@ func BodyLimitWithConfig(config BodyLimitConfig) echo.MiddlewareFunc { } limit, err := bytes.Parse(config.Limit) - if err != nil { + if err != nil || limit < 0 { panic(fmt.Errorf("echo: invalid body-limit=%s", config.Limit)) } config.limit = limit - pool := limitedReaderPool(config) return func(next echo.HandlerFunc) echo.HandlerFunc { return func(c echo.Context) error { @@ -78,13 +76,7 @@ func BodyLimitWithConfig(config BodyLimitConfig) echo.MiddlewareFunc { } // 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 + req.Body = &limitedReader{BodyLimitConfig: config, reader: req.Body} return next(c) } @@ -92,33 +84,35 @@ func BodyLimitWithConfig(config BodyLimitConfig) echo.MiddlewareFunc { } func (r *limitedReader) Read(b []byte) (n int, err error) { - if r.limit > 0 && r.read > r.limit { - return 0, echo.ErrStatusRequestEntityTooLarge + // A zero limit has historically disabled the read limit in v4. + if r.limit == 0 { + return r.reader.Read(b) + } + if r.err != nil { + return 0, r.err + } + if len(b) == 0 { + return 0, nil } + remaining := r.limit - r.read + // Read at most one byte beyond the limit to distinguish an exact-size body + // from an oversized one without exposing the extra byte to the caller. + // Unlike http.MaxBytesReader, this reader does not signal net/http to close + // the connection after the limit is exceeded. + if int64(len(b)) > remaining { + b = b[:remaining+1] + } n, err = r.reader.Read(b) - r.read += int64(n) - - if r.limit > 0 && r.read > r.limit { - return n, echo.ErrStatusRequestEntityTooLarge + if int64(n) > remaining { + r.read = r.limit + r.err = echo.ErrStatusRequestEntityTooLarge + return int(remaining), r.err } - + r.read += int64(n) return n, err } func (r *limitedReader) Close() error { return r.reader.Close() } - -func (r *limitedReader) Reset(reader io.ReadCloser) { - r.reader = reader - r.read = 0 -} - -func limitedReaderPool(c BodyLimitConfig) sync.Pool { - return sync.Pool{ - New: func() interface{} { - return &limitedReader{BodyLimitConfig: c} - }, - } -} diff --git a/middleware/body_limit_test.go b/middleware/body_limit_test.go index 8ad9004c1..16d2de250 100644 --- a/middleware/body_limit_test.go +++ b/middleware/body_limit_test.go @@ -5,7 +5,11 @@ package middleware import ( "bytes" + "errors" + "fmt" "io" + "math" + "mime/multipart" "net/http" "net/http/httptest" "testing" @@ -75,14 +79,143 @@ func TestBodyLimitReader(t *testing.T) { he := err.(*echo.HTTPError) assert.Equal(t, http.StatusRequestEntityTooLarge, he.Code) - // reset reader and read two bytes must succeed + // A new request gets a fresh reader. 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 TestBodyLimitReaderDoesNotExposeExtraBytes(t *testing.T) { + for _, chunkSize := range []int{1, 64} { + t.Run(fmt.Sprint(chunkSize), func(t *testing.T) { + reader := &limitedReader{ + BodyLimitConfig: BodyLimitConfig{limit: 5}, + reader: io.NopCloser(bytes.NewReader([]byte("123456789"))), + } + buf := make([]byte, chunkSize) + var got []byte + for { + n, err := reader.Read(buf) + got = append(got, buf[:n]...) + if err != nil { + assert.Equal(t, []byte("12345"), got) + assertBodyLimitError(t, err) + break + } + } + n, err := reader.Read(buf) + assert.Zero(t, n) + assertBodyLimitError(t, err) + }) + } +} + +func TestBodyLimitReaderBoundaryAndLargeLimit(t *testing.T) { + for _, limit := range []int64{5, math.MaxInt64} { + reader := &limitedReader{ + BodyLimitConfig: BodyLimitConfig{limit: limit}, + reader: io.NopCloser(bytes.NewReader([]byte("12345"))), + } + got, err := io.ReadAll(reader) + assert.NoError(t, err) + assert.Equal(t, "12345", string(got)) + } +} + +func TestBodyLimitReaderKeepsOrdinaryReadErrors(t *testing.T) { + broken := errors.New("source failed") + reader := &limitedReader{ + BodyLimitConfig: BodyLimitConfig{limit: 5}, + reader: io.NopCloser(&errorOnceReader{err: broken}), + } + buf := make([]byte, 2) + _, err := reader.Read(buf) + assert.ErrorIs(t, err, broken) + n, err := reader.Read(buf) + assert.NoError(t, err) + assert.Equal(t, 1, n) + assert.Equal(t, byte('x'), buf[0]) +} + +type errorOnceReader struct { + err error +} + +func (r *errorOnceReader) Read(b []byte) (int, error) { + if r.err != nil { + err := r.err + r.err = nil + return 0, err + } + return copy(b, "x"), nil +} + +func TestBodyLimitReaderRemainsAttachedToRequest(t *testing.T) { + e := echo.New() + req := httptest.NewRequest(http.MethodPost, "/", bytes.NewBufferString("123456")) + req.ContentLength = -1 + c := e.NewContext(req, httptest.NewRecorder()) + err := BodyLimit("5B")(func(c echo.Context) error { return nil })(c) + assert.NoError(t, err) + got, err := io.ReadAll(req.Body) + assert.Equal(t, "12345", string(got)) + assertBodyLimitError(t, err) +} + +func TestBodyLimitBindRejectsOversizeBodies(t *testing.T) { + var multipartBody bytes.Buffer + writer := multipart.NewWriter(&multipartBody) + if err := writer.WriteField("x", "12345"); err != nil { + t.Fatal(err) + } + if err := writer.Close(); err != nil { + t.Fatal(err) + } + + tests := []struct { + name, contentType string + body []byte + }{ + {"JSON", echo.MIMEApplicationJSON, []byte(`{"x":"12345"}`)}, + {"XML", echo.MIMEApplicationXML, []byte(`12345`)}, + {"form", echo.MIMEApplicationForm, []byte(`x=12345`)}, + {"multipart", writer.FormDataContentType(), multipartBody.Bytes()}, + } + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + var bindErr error + e := echo.New() + e.Use(BodyLimit("5B")) + e.POST("/", func(c echo.Context) error { + var payload struct { + X string `json:"x" xml:",chardata" form:"x"` + } + bindErr = c.Bind(&payload) + return bindErr + }) + req := httptest.NewRequest(http.MethodPost, "/", bytes.NewReader(tc.body)) + req.ContentLength = -1 + req.Header.Set(echo.HeaderContentType, tc.contentType) + rec := httptest.NewRecorder() + e.ServeHTTP(rec, req) + assert.Equal(t, http.StatusRequestEntityTooLarge, rec.Code, rec.Body.String()) + // Bind itself must return the 413, not a 400 that only the error handler turns into 413. + assertBodyLimitError(t, bindErr) + }) + } +} + +func assertBodyLimitError(t *testing.T, err error) { + t.Helper() + var httpError *echo.HTTPError + if !errors.As(err, &httpError) { + t.Fatalf("expected HTTPError, got %v", err) + } + assert.Equal(t, http.StatusRequestEntityTooLarge, httpError.Code) +} + func TestBodyLimitWithConfig_Skipper(t *testing.T) { e := echo.New() h := func(c echo.Context) error { @@ -170,6 +303,11 @@ func TestBodyLimit_panicOnInvalidLimit(t *testing.T) { "echo: invalid body-limit=", func() { BodyLimit("") }, ) + assert.PanicsWithError( + t, + "echo: invalid body-limit=-1B", + func() { BodyLimit("-1B") }, + ) } func TestBodyLimit_Middleware_BodyRestoration(t *testing.T) {