Skip to content

Commit cb529b4

Browse files
committed
fix(archive): limit total decompressed size to prevent disk exhaustion
1 parent aead76e commit cb529b4

13 files changed

Lines changed: 125 additions & 28 deletions

File tree

internal/archive/archives/archives.go

Lines changed: 5 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -10,7 +10,9 @@ import (
1010
"strings"
1111

1212
"github.com/alist-org/alist/v3/internal/archive/tool"
13+
"github.com/alist-org/alist/v3/internal/conf"
1314
"github.com/alist-org/alist/v3/internal/model"
15+
"github.com/alist-org/alist/v3/internal/setting"
1416
"github.com/alist-org/alist/v3/internal/stream"
1517
"github.com/alist-org/alist/v3/pkg/utils"
1618
)
@@ -96,6 +98,7 @@ func (Archives) Decompress(ss []*stream.SeekableStream, outputPath string, args
9698
if err != nil {
9799
return err
98100
}
101+
limiter := tool.NewSizeLimiter(int64(setting.GetInt(conf.MaxExtractSize, 0)) << 30)
99102
isDir := false
100103
path := strings.TrimPrefix(args.InnerPath, "/")
101104
if path == "" {
@@ -148,7 +151,7 @@ func (Archives) Decompress(ss []*stream.SeekableStream, outputPath string, args
148151
if err := os.MkdirAll(filepath.Dir(dstPath), 0700); err != nil {
149152
return err
150153
}
151-
return decompress(fsys, p, dstPath, func(_ float64) {})
154+
return decompress(fsys, p, dstPath, func(_ float64) {}, limiter)
152155
})
153156
} else {
154157
entryName := stdpath.Base(path)
@@ -159,7 +162,7 @@ func (Archives) Decompress(ss []*stream.SeekableStream, outputPath string, args
159162
if err = os.MkdirAll(filepath.Dir(dstPath), 0700); err != nil {
160163
return err
161164
}
162-
err = decompress(fsys, path, dstPath, up)
165+
err = decompress(fsys, path, dstPath, up, limiter)
163166
}
164167
return filterPassword(err)
165168
}

internal/archive/archives/utils.go

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -60,7 +60,7 @@ func filterPassword(err error) error {
6060
return err
6161
}
6262

63-
func decompress(fsys fs2.FS, filePath, dstPath string, up model.UpdateProgress) error {
63+
func decompress(fsys fs2.FS, filePath, dstPath string, up model.UpdateProgress, limiter *tool.SizeLimiter) error {
6464
rc, err := fsys.Open(filePath)
6565
if err != nil {
6666
return err
@@ -78,7 +78,7 @@ func decompress(fsys fs2.FS, filePath, dstPath string, up model.UpdateProgress)
7878
return err
7979
}
8080
defer f.Close()
81-
_, err = utils.CopyWithBuffer(f, &stream.ReaderUpdatingProgress{
81+
_, err = utils.CopyWithBuffer(limiter.WrapWriter(f), &stream.ReaderUpdatingProgress{
8282
Reader: &stream.SimpleReaderWithSize{
8383
Reader: rc,
8484
Size: stat.Size(),

internal/archive/iso9660/iso9660.go

Lines changed: 5 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -5,8 +5,10 @@ import (
55
"os"
66

77
"github.com/alist-org/alist/v3/internal/archive/tool"
8+
"github.com/alist-org/alist/v3/internal/conf"
89
"github.com/alist-org/alist/v3/internal/errs"
910
"github.com/alist-org/alist/v3/internal/model"
11+
"github.com/alist-org/alist/v3/internal/setting"
1012
"github.com/alist-org/alist/v3/internal/stream"
1113
"github.com/kdomanski/iso9660"
1214
)
@@ -76,6 +78,7 @@ func (ISO9660) Decompress(ss []*stream.SeekableStream, outputPath string, args m
7678
if err != nil {
7779
return err
7880
}
81+
limiter := tool.NewSizeLimiter(int64(setting.GetInt(conf.MaxExtractSize, 0)) << 30)
7982
if obj.IsDir() {
8083
if args.InnerPath != "/" {
8184
outputPath, err = tool.SecureJoin(outputPath, obj.Name())
@@ -88,10 +91,10 @@ func (ISO9660) Decompress(ss []*stream.SeekableStream, outputPath string, args m
8891
}
8992
var children []*iso9660.File
9093
if children, err = obj.GetChildren(); err == nil {
91-
err = decompressAll(children, outputPath)
94+
err = decompressAll(children, outputPath, limiter)
9295
}
9396
} else {
94-
err = decompress(obj, outputPath, up)
97+
err = decompress(obj, outputPath, up, limiter)
9598
}
9699
return err
97100
}

internal/archive/iso9660/utils.go

Lines changed: 7 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -63,11 +63,11 @@ func toModelObj(file *iso9660.File) model.Obj {
6363
}
6464
}
6565

66-
func decompress(f *iso9660.File, path string, up model.UpdateProgress) error {
67-
return decompressEntry(f.Reader(), f.Size(), path, f.Name(), up)
66+
func decompress(f *iso9660.File, path string, up model.UpdateProgress, limiter *tool.SizeLimiter) error {
67+
return decompressEntry(f.Reader(), f.Size(), path, f.Name(), up, limiter)
6868
}
6969

70-
func decompressEntry(reader io.Reader, size int64, path, entryName string, up model.UpdateProgress) error {
70+
func decompressEntry(reader io.Reader, size int64, path, entryName string, up model.UpdateProgress, limiter *tool.SizeLimiter) error {
7171
dstPath, err := tool.SecureJoin(path, entryName)
7272
if err != nil {
7373
return err
@@ -80,7 +80,7 @@ func decompressEntry(reader io.Reader, size int64, path, entryName string, up mo
8080
return err
8181
}
8282
defer file.Close()
83-
_, err = utils.CopyWithBuffer(file, &stream.ReaderUpdatingProgress{
83+
_, err = utils.CopyWithBuffer(limiter.WrapWriter(file), &stream.ReaderUpdatingProgress{
8484
Reader: &stream.SimpleReaderWithSize{
8585
Reader: reader,
8686
Size: size,
@@ -90,7 +90,7 @@ func decompressEntry(reader io.Reader, size int64, path, entryName string, up mo
9090
return err
9191
}
9292

93-
func decompressAll(children []*iso9660.File, path string) error {
93+
func decompressAll(children []*iso9660.File, path string, limiter *tool.SizeLimiter) error {
9494
for _, child := range children {
9595
if child.IsDir() {
9696
nextChildren, err := child.GetChildren()
@@ -104,11 +104,11 @@ func decompressAll(children []*iso9660.File, path string) error {
104104
if err = os.MkdirAll(nextPath, 0700); err != nil {
105105
return err
106106
}
107-
if err = decompressAll(nextChildren, nextPath); err != nil {
107+
if err = decompressAll(nextChildren, nextPath, limiter); err != nil {
108108
return err
109109
}
110110
} else {
111-
if err := decompress(child, path, func(_ float64) {}); err != nil {
111+
if err := decompress(child, path, func(_ float64) {}, limiter); err != nil {
112112
return err
113113
}
114114
}

internal/archive/rardecode/rardecode.go

Lines changed: 6 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -9,8 +9,10 @@ import (
99
"strings"
1010

1111
"github.com/alist-org/alist/v3/internal/archive/tool"
12+
"github.com/alist-org/alist/v3/internal/conf"
1213
"github.com/alist-org/alist/v3/internal/errs"
1314
"github.com/alist-org/alist/v3/internal/model"
15+
"github.com/alist-org/alist/v3/internal/setting"
1416
"github.com/alist-org/alist/v3/internal/stream"
1517
"github.com/nwaples/rardecode/v2"
1618
)
@@ -74,6 +76,7 @@ func (RarDecoder) Decompress(ss []*stream.SeekableStream, outputPath string, arg
7476
if err != nil {
7577
return err
7678
}
79+
limiter := tool.NewSizeLimiter(int64(setting.GetInt(conf.MaxExtractSize, 0)) << 30)
7780
if args.InnerPath == "/" {
7881
for {
7982
var header *rardecode.FileHeader
@@ -92,7 +95,7 @@ func (RarDecoder) Decompress(ss []*stream.SeekableStream, outputPath string, arg
9295
if e != nil {
9396
return e
9497
}
95-
err = decompress(reader, header, dstPath)
98+
err = decompress(reader, header, dstPath, limiter)
9699
if err != nil {
97100
return err
98101
}
@@ -139,7 +142,7 @@ func (RarDecoder) Decompress(ss []*stream.SeekableStream, outputPath string, arg
139142
if err = os.MkdirAll(filepath.Dir(dstPath), 0700); err != nil {
140143
return err
141144
}
142-
err = _decompress(reader, header, dstPath, up)
145+
err = _decompress(reader, header, dstPath, up, limiter)
143146
if err != nil {
144147
return err
145148
}
@@ -164,7 +167,7 @@ func (RarDecoder) Decompress(ss []*stream.SeekableStream, outputPath string, arg
164167
if e != nil {
165168
return e
166169
}
167-
err = decompress(reader, header, dstPath)
170+
err = decompress(reader, header, dstPath, limiter)
168171
if err != nil {
169172
return err
170173
}

internal/archive/rardecode/utils.go

Lines changed: 4 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -181,7 +181,7 @@ func getReader(ss []*stream.SeekableStream, password string) (*rardecode.Reader,
181181
return &rc.Reader, nil
182182
}
183183

184-
func decompress(reader *rardecode.Reader, header *rardecode.FileHeader, dstPath string) error {
184+
func decompress(reader *rardecode.Reader, header *rardecode.FileHeader, dstPath string, limiter *tool.SizeLimiter) error {
185185
if header.IsDir {
186186
return os.MkdirAll(dstPath, 0700)
187187
}
@@ -191,16 +191,16 @@ func decompress(reader *rardecode.Reader, header *rardecode.FileHeader, dstPath
191191
if err := os.MkdirAll(filepath.Dir(dstPath), 0700); err != nil {
192192
return err
193193
}
194-
return _decompress(reader, header, dstPath, func(_ float64) {})
194+
return _decompress(reader, header, dstPath, func(_ float64) {}, limiter)
195195
}
196196

197-
func _decompress(reader *rardecode.Reader, header *rardecode.FileHeader, dstPath string, up model.UpdateProgress) error {
197+
func _decompress(reader *rardecode.Reader, header *rardecode.FileHeader, dstPath string, up model.UpdateProgress, limiter *tool.SizeLimiter) error {
198198
f, err := os.OpenFile(dstPath, os.O_WRONLY|os.O_CREATE|os.O_EXCL, 0600)
199199
if err != nil {
200200
return err
201201
}
202202
defer func() { _ = f.Close() }()
203-
_, err = io.Copy(f, &stream.ReaderUpdatingProgress{
203+
_, err = io.Copy(limiter.WrapWriter(f), &stream.ReaderUpdatingProgress{
204204
Reader: &stream.SimpleReaderWithSize{
205205
Reader: reader,
206206
Size: header.UnPackedSize,

internal/archive/sevenzip/sevenzip.go

Lines changed: 4 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -5,8 +5,10 @@ import (
55
"strings"
66

77
"github.com/alist-org/alist/v3/internal/archive/tool"
8+
"github.com/alist-org/alist/v3/internal/conf"
89
"github.com/alist-org/alist/v3/internal/errs"
910
"github.com/alist-org/alist/v3/internal/model"
11+
"github.com/alist-org/alist/v3/internal/setting"
1012
"github.com/alist-org/alist/v3/internal/stream"
1113
)
1214

@@ -62,7 +64,8 @@ func (SevenZip) Decompress(ss []*stream.SeekableStream, outputPath string, args
6264
if err != nil {
6365
return err
6466
}
65-
return tool.DecompressFromFolderTraversal(&WrapReader{Reader: reader}, outputPath, args, up)
67+
limiter := tool.NewSizeLimiter(int64(setting.GetInt(conf.MaxExtractSize, 0)) << 30)
68+
return tool.DecompressFromFolderTraversal(&WrapReader{Reader: reader}, outputPath, args, up, limiter)
6669
}
6770

6871
var _ tool.Tool = (*SevenZip)(nil)

internal/archive/tool/helper.go

Lines changed: 6 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -115,7 +115,7 @@ type WrapFileInfo struct {
115115
model.Obj
116116
}
117117

118-
func DecompressFromFolderTraversal(r ArchiveReader, outputPath string, args model.ArchiveInnerArgs, up model.UpdateProgress) error {
118+
func DecompressFromFolderTraversal(r ArchiveReader, outputPath string, args model.ArchiveInnerArgs, up model.UpdateProgress, limiter *SizeLimiter) error {
119119
var err error
120120
files := r.Files()
121121
if args.InnerPath == "/" {
@@ -144,7 +144,7 @@ func DecompressFromFolderTraversal(r ArchiveReader, outputPath string, args mode
144144
if err = os.MkdirAll(filepath.Dir(dstPath), 0700); err != nil {
145145
return err
146146
}
147-
err = _decompress(file, dstPath, args.Password, func(_ float64) {})
147+
err = _decompress(file, dstPath, args.Password, func(_ float64) {}, limiter)
148148
if err != nil {
149149
return err
150150
}
@@ -183,7 +183,7 @@ func DecompressFromFolderTraversal(r ArchiveReader, outputPath string, args mode
183183
if err = os.MkdirAll(filepath.Dir(dstPath), 0700); err != nil {
184184
return err
185185
}
186-
err = _decompress(file, dstPath, args.Password, up)
186+
err = _decompress(file, dstPath, args.Password, up, limiter)
187187
if err != nil {
188188
return err
189189
}
@@ -227,7 +227,7 @@ func DecompressFromFolderTraversal(r ArchiveReader, outputPath string, args mode
227227
if err = os.MkdirAll(filepath.Dir(dstPath), 0700); err != nil {
228228
return err
229229
}
230-
err = _decompress(file, dstPath, args.Password, func(_ float64) {})
230+
err = _decompress(file, dstPath, args.Password, func(_ float64) {}, limiter)
231231
if err != nil {
232232
return err
233233
}
@@ -237,7 +237,7 @@ func DecompressFromFolderTraversal(r ArchiveReader, outputPath string, args mode
237237
return nil
238238
}
239239

240-
func _decompress(file SubFile, dstPath, password string, up model.UpdateProgress) error {
240+
func _decompress(file SubFile, dstPath, password string, up model.UpdateProgress, limiter *SizeLimiter) error {
241241
if encrypt, ok := file.(CanEncryptSubFile); ok && encrypt.IsEncrypted() {
242242
encrypt.SetPassword(password)
243243
}
@@ -251,7 +251,7 @@ func _decompress(file SubFile, dstPath, password string, up model.UpdateProgress
251251
return err
252252
}
253253
defer func() { _ = f.Close() }()
254-
_, err = io.Copy(f, &stream.ReaderUpdatingProgress{
254+
_, err = io.Copy(limiter.WrapWriter(f), &stream.ReaderUpdatingProgress{
255255
Reader: &stream.SimpleReaderWithSize{
256256
Reader: rc,
257257
Size: file.FileInfo().Size(),

internal/archive/tool/limiter.go

Lines changed: 42 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,42 @@
1+
package tool
2+
3+
import (
4+
"errors"
5+
"io"
6+
"sync/atomic"
7+
)
8+
9+
var ErrExtractSizeExceeded = errors.New("total size of decompressed files exceeds the limit")
10+
11+
// SizeLimiter limits the total bytes written by one decompress task.
12+
// A non-positive max means no limit.
13+
type SizeLimiter struct {
14+
remain int64
15+
limited bool
16+
}
17+
18+
func NewSizeLimiter(max int64) *SizeLimiter {
19+
if max <= 0 {
20+
return &SizeLimiter{}
21+
}
22+
return &SizeLimiter{remain: max, limited: true}
23+
}
24+
25+
func (l *SizeLimiter) WrapWriter(w io.Writer) io.Writer {
26+
if l == nil || !l.limited {
27+
return w
28+
}
29+
return &limitedWriter{w: w, limiter: l}
30+
}
31+
32+
type limitedWriter struct {
33+
w io.Writer
34+
limiter *SizeLimiter
35+
}
36+
37+
func (lw *limitedWriter) Write(p []byte) (int, error) {
38+
if atomic.AddInt64(&lw.limiter.remain, -int64(len(p))) < 0 {
39+
return 0, ErrExtractSizeExceeded
40+
}
41+
return lw.w.Write(p)
42+
}
Lines changed: 38 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,38 @@
1+
package tool
2+
3+
import (
4+
"bytes"
5+
"errors"
6+
"io"
7+
"strings"
8+
"testing"
9+
)
10+
11+
func TestSizeLimiterExceeded(t *testing.T) {
12+
l := NewSizeLimiter(10)
13+
var buf bytes.Buffer
14+
_, err := io.Copy(l.WrapWriter(&buf), strings.NewReader("0123456789abcdef"))
15+
if !errors.Is(err, ErrExtractSizeExceeded) {
16+
t.Fatalf("expected ErrExtractSizeExceeded, got %v", err)
17+
}
18+
}
19+
20+
func TestSizeLimiterSharedAcrossWriters(t *testing.T) {
21+
l := NewSizeLimiter(10)
22+
var a, b bytes.Buffer
23+
if _, err := l.WrapWriter(&a).Write([]byte("123456")); err != nil {
24+
t.Fatalf("first write should pass, got %v", err)
25+
}
26+
if _, err := l.WrapWriter(&b).Write([]byte("123456")); !errors.Is(err, ErrExtractSizeExceeded) {
27+
t.Fatalf("expected ErrExtractSizeExceeded, got %v", err)
28+
}
29+
}
30+
31+
func TestSizeLimiterUnlimited(t *testing.T) {
32+
l := NewSizeLimiter(0)
33+
var buf bytes.Buffer
34+
n, err := io.Copy(l.WrapWriter(&buf), strings.NewReader("0123456789abcdef"))
35+
if err != nil || n != 16 {
36+
t.Fatalf("unlimited limiter should pass all data, n=%d err=%v", n, err)
37+
}
38+
}

0 commit comments

Comments
 (0)