From 7ff7e0f043f036cba9511100e3b41bf0cff5e56a Mon Sep 17 00:00:00 2001 From: Gani Georgiev Date: Sun, 6 Sep 2026 19:51:06 +0300 Subject: [PATCH] handled io.ReadAll error with exact payload size --- apis/middlewares_body_limit.go | 52 +++++++++++++++++++---------- apis/middlewares_body_limit_test.go | 18 ++++++++++ 2 files changed, 52 insertions(+), 18 deletions(-) diff --git a/apis/middlewares_body_limit.go b/apis/middlewares_body_limit.go index 70b3fb5e..1c456156 100644 --- a/apis/middlewares_body_limit.go +++ b/apis/middlewares_body_limit.go @@ -84,57 +84,73 @@ func applyBodyLimit(e *core.RequestEvent, limitBytes int64) error { } // replace the request body - e.Request.Body = newLimitedReader(e.Request.Body, limitBytes) + e.Request.Body = newMaxBytesReader(e.Request.Body, limitBytes) return nil } -func newLimitedReader(body io.ReadCloser, limitBytes int64) *limitedReader { - return &limitedReader{ +func newMaxBytesReader(body io.ReadCloser, limitBytes int64) *maxBytesReader { + return &maxBytesReader{ ReadCloser: body, limit: limitBytes, remaining: limitBytes, } } -type limitedReader struct { +// maxBytesReader is very similar to the http.MaxBytesReader but support +// rereads and doesn't try to prematurely close the related response +// to allow consequent middlewares to operate correctly. +type maxBytesReader struct { io.ReadCloser limit int64 remaining int64 + stickyErr error } -func (r *limitedReader) Read(b []byte) (int, error) { - if r.remaining <= 0 { - return 0, ErrRequestEntityTooLarge +func (r *maxBytesReader) Read(b []byte) (int, error) { + if r.stickyErr != nil { + return 0, r.stickyErr } - if int64(len(b)) > r.remaining { - b = b[0:r.remaining] + if len(b) == 0 { + return 0, nil + } + + // if possible no need to read the entire chunk since + // remaining+1 is enough to determine whether it exceed the limit + if int64(len(b))-1 > r.remaining { + b = b[:r.remaining+1] } n, err := r.ReadCloser.Read(b) - r.remaining -= int64(n) + if int64(n) <= r.remaining { + r.remaining -= int64(n) + r.stickyErr = err + return n, err + } - return n, err + n = int(r.remaining) + + r.remaining = 0 + r.stickyErr = ErrRequestEntityTooLarge + + return n, r.stickyErr } // explicit casts to ensure that the main struct methods will be invoked // (extra precautions in case of nested interface wrapping erasure) // --- -func (r *limitedReader) Reread() { +func (r *maxBytesReader) Reread() { rereader, ok := r.ReadCloser.(router.Rereader) if ok { rereader.Reread() r.remaining = r.limit + r.stickyErr = nil } } -func (r *limitedReader) Close() error { - closer, ok := r.ReadCloser.(io.Closer) - if ok { - return closer.Close() - } - return nil +func (r *maxBytesReader) Close() error { + return r.ReadCloser.Close() } diff --git a/apis/middlewares_body_limit_test.go b/apis/middlewares_body_limit_test.go index abfd48b0..f4ac9ae7 100644 --- a/apis/middlewares_body_limit_test.go +++ b/apis/middlewares_body_limit_test.go @@ -1,6 +1,7 @@ package apis_test import ( + "io" "net/http" "net/http/httptest" "strings" @@ -39,6 +40,15 @@ func TestBodyLimitMiddleware(t *testing.T) { pbRouter.POST("/a", testHandler) // default global BodyLimit check pbRouter.POST("/b", testHandler).Bind(apis.BodyLimit(customLimit)) + pbRouter.POST("/iof", func(e *core.RequestEvent) error { + // ensure that normal io methods still operate correctly + b, err := io.ReadAll(e.Request.Body) + if err != nil { + return err + } + + return e.String(http.StatusOK, string(b)) + }).Bind(apis.BodyLimit(customLimit)) mux, err := pbRouter.BuildMux() if err != nil { @@ -109,6 +119,14 @@ func TestBodyLimitMiddleware(t *testing.T) { true, http.StatusRequestEntityTooLarge, }, + // --- + { + "io.ReadAll io.EOF exact limit check", + "/iof", + `"` + strings.Repeat("a", customLimit-2) + `"`, + true, + http.StatusOK, + }, } for _, s := range scenarios {