diff --git a/middleware/decompress.go b/middleware/decompress.go index 52febc680..355455ce6 100644 --- a/middleware/decompress.go +++ b/middleware/decompress.go @@ -128,7 +128,6 @@ func (config DecompressConfig) ToMiddleware() (echo.MiddlewareFunc, error) { }, nil } - // isGzipContentEncoding reports whether Content-Encoding is gzip. // Content codings are case-insensitive per RFC 9110 ยง8.4.1. func isGzipContentEncoding(v string) bool { @@ -143,10 +142,22 @@ type limitedGzipReader struct { } func (r *limitedGzipReader) Read(p []byte) (n int, err error) { - if r.remaining <= 0 { - // Limit exceeded - return 413 error + if len(p) == 0 { + return 0, nil + } + if r.remaining < 0 { return 0, echo.ErrStatusRequestEntityTooLarge } + if r.remaining == 0 { + // Reaching the limit is valid if there is no more decompressed data. + var probe [1]byte + n, err = r.Reader.Read(probe[:]) + if n > 0 { + r.remaining = -1 + return 0, echo.ErrStatusRequestEntityTooLarge + } + return 0, err + } // Limit the read to remaining bytes if int64(len(p)) > r.remaining { diff --git a/middleware/decompress_test.go b/middleware/decompress_test.go index fadfa7fac..b154dd5b2 100644 --- a/middleware/decompress_test.go +++ b/middleware/decompress_test.go @@ -320,6 +320,52 @@ func TestDecompress_AtExactLimit(t *testing.T) { assert.Equal(t, exactBody, rec.Body.String()) } +func TestDecompress_ReadAfterExactLimit(t *testing.T) { + e := echo.New() + body := strings.Repeat("B", 1024) + gz, err := gzipString(body) + assert.NoError(t, err) + req := httptest.NewRequest(http.MethodPost, "/", bytes.NewReader(gz)) + req.Header.Set(echo.HeaderContentEncoding, GZIPEncoding) + c := e.NewContext(req, httptest.NewRecorder()) + h := DecompressWithConfig(DecompressConfig{MaxDecompressedSize: 1024}) + + err = h(func(c *echo.Context) error { + buf := make([]byte, 1024) + n, readErr := io.ReadFull(c.Request().Body, buf) + assert.NoError(t, readErr) + assert.Equal(t, body, string(buf[:n])) + n, readErr = c.Request().Body.Read(nil) + assert.Zero(t, n) + assert.NoError(t, readErr) + n, readErr = c.Request().Body.Read(buf) + assert.Zero(t, n) + assert.ErrorIs(t, readErr, io.EOF) + return nil + })(c) + assert.NoError(t, err) +} + +func TestLimitedGzipReader_ExceededLimitStaysExceeded(t *testing.T) { + gz, err := gzipString("AB") + assert.NoError(t, err) + reader, err := gzip.NewReader(bytes.NewReader(gz)) + assert.NoError(t, err) + defer reader.Close() + r := &limitedGzipReader{Reader: reader, remaining: 1, limit: 1} + buf := make([]byte, 1) + n, err := io.ReadFull(r, buf) + assert.NoError(t, err) + assert.Equal(t, 1, n) + assert.Equal(t, "A", string(buf)) + + for i := 0; i < 2; i++ { + n, err = r.Read(buf) + assert.Zero(t, n) + assert.ErrorIs(t, err, echo.ErrStatusRequestEntityTooLarge) + } +} + func TestDecompress_ZipBomb(t *testing.T) { e := echo.New() // Create highly compressed data that expands to 2MB