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
12 changes: 12 additions & 0 deletions bind.go
Original file line number Diff line number Diff line change
Expand Up @@ -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 {
Expand All @@ -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 {
Expand All @@ -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 {
Expand Down
56 changes: 25 additions & 31 deletions middleware/body_limit.go
Original file line number Diff line number Diff line change
Expand Up @@ -6,8 +6,6 @@ package middleware
import (
"fmt"
"io"
"net/http"
"sync"

"github.com/labstack/echo/v4"
"github.com/labstack/gommon/bytes"
Expand All @@ -28,6 +26,7 @@ type limitedReader struct {
BodyLimitConfig
reader io.ReadCloser
read int64
err error
}

// DefaultBodyLimitConfig is the default BodyLimit middleware config.
Expand Down Expand Up @@ -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 {
Expand All @@ -78,47 +76,43 @@ 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)
}
}
}

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}
},
}
}
142 changes: 140 additions & 2 deletions middleware/body_limit_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -5,7 +5,11 @@ package middleware

import (
"bytes"
"errors"
"fmt"
"io"
"math"
"mime/multipart"
"net/http"
"net/http/httptest"
"testing"
Expand Down Expand Up @@ -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(`<x>12345</x>`)},
{"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 {
Expand Down Expand Up @@ -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) {
Expand Down
Loading