diff --git a/excelize_test.go b/excelize_test.go index 18baef1448..403dd0d4ae 100644 --- a/excelize_test.go +++ b/excelize_test.go @@ -303,6 +303,8 @@ func TestOpenReader(t *testing.T) { } // Test open spreadsheet with unzip size limit + _, err = OpenFile(filepath.Join("test", "Book1.xlsx"), Options{UnzipSizeLimit: -1}) + assert.EqualError(t, err, newUnzipSizeLimitError(-1).Error()) _, err = OpenFile(filepath.Join("test", "Book1.xlsx"), Options{UnzipSizeLimit: 100}) assert.EqualError(t, err, newUnzipSizeLimitError(100).Error()) diff --git a/lib.go b/lib.go index 5c49705b94..3b3a8eed5c 100644 --- a/lib.go +++ b/lib.go @@ -28,6 +28,18 @@ import ( "unicode/utf16" ) +// checkFileSize checks if the file size and unzip size exceed the limit set in +// options. +func (f *File) checkFileSize(fileSize, unzipSize int64) error { + if f.options.UnzipSizeLimit < 0 || uint64(f.options.UnzipSizeLimit) < uint64(fileSize) || fileSize < 0 { + return newUnzipSizeLimitError(f.options.UnzipSizeLimit) + } + if unzipSize > f.options.UnzipSizeLimit { + return newUnzipSizeLimitError(f.options.UnzipSizeLimit) + } + return nil +} + // ReadZipReader extract spreadsheet with given options. func (f *File) ReadZipReader(r *zip.Reader) (map[string][]byte, int, error) { var ( @@ -43,8 +55,8 @@ func (f *File) ReadZipReader(r *zip.Reader) (map[string][]byte, int, error) { for _, v := range r.File { fileSize := v.FileInfo().Size() unzipSize += fileSize - if unzipSize > f.options.UnzipSizeLimit { - return fileList, worksheets, newUnzipSizeLimitError(f.options.UnzipSizeLimit) + if err := f.checkFileSize(fileSize, unzipSize); err != nil { + return fileList, worksheets, err } fileName := strings.ReplaceAll(v.Name, "\\", "/") if partName, ok := docPart[strings.ToLower(fileName)]; ok { diff --git a/lib_test.go b/lib_test.go index 225e24b668..9a804f6e24 100644 --- a/lib_test.go +++ b/lib_test.go @@ -422,3 +422,11 @@ func TestFloat2Frac(t *testing.T) { assert.Equal(t, "9999/10000", strings.Trim(floatToFraction(0.9999, 10, 10), " ")) assert.Equal(t, "954888175898973913/351283728530932463", floatToFraction(math.E, 1, 18)) } + +func TestCheckFileSize(t *testing.T) { + f := NewFile() + assert.NoError(t, f.checkFileSize(1, 1)) + assert.EqualError(t, f.checkFileSize(UnzipSizeLimit+1, 1), newUnzipSizeLimitError(UnzipSizeLimit).Error()) + assert.EqualError(t, f.checkFileSize(1, UnzipSizeLimit+1), newUnzipSizeLimitError(UnzipSizeLimit).Error()) + assert.EqualError(t, f.checkFileSize(-1, 1), newUnzipSizeLimitError(UnzipSizeLimit).Error()) +}