diff --git a/drivers/alias/driver.go b/drivers/alias/driver.go index d69d6cf502..40c9cc564b 100644 --- a/drivers/alias/driver.go +++ b/drivers/alias/driver.go @@ -259,24 +259,16 @@ func (d *Alias) Link(ctx context.Context, file model.Obj, args model.LinkArgs) ( _ = link.Close() continue } - l := &model.Link{ - URL: link.URL, - Header: link.Header, - RangeReader: link.RangeReader, - Concurrency: link.Concurrency, - PartSize: link.PartSize, - ContentLength: link.ContentLength, - } if d.DownloadConcurrency > 0 { - l.Concurrency = d.DownloadConcurrency + link.Concurrency = d.DownloadConcurrency } if d.DownloadPartSize > 0 { - l.PartSize = d.DownloadPartSize * utils.KB + link.PartSize = d.DownloadPartSize * utils.KB } - if l.ContentLength == 0 { - l.ContentLength = fi.GetSize() + if link.ContentLength == 0 { + link.ContentLength = fi.GetSize() } - rr, err := stream.GetRangeReaderFromLink(l.ContentLength, l) + rr, err := stream.GetRangeReaderFromLink(link.ContentLength, link) if err != nil { _ = link.Close() continue @@ -327,20 +319,19 @@ func (d *Alias) Link(ctx context.Context, file model.Obj, args model.LinkArgs) ( if err != nil { return nil, err } - resultLink := link.Clone() // 复制一份,避免修改到原始link if args.Redirect { - return resultLink, nil + return link, nil } if d.DownloadConcurrency > 0 { - resultLink.Concurrency = d.DownloadConcurrency + link.Concurrency = d.DownloadConcurrency } if d.DownloadPartSize > 0 { - resultLink.PartSize = d.DownloadPartSize * utils.KB + link.PartSize = d.DownloadPartSize * utils.KB } - if resultLink.ContentLength == 0 { - resultLink.ContentLength = fi.GetSize() + if link.ContentLength == 0 { + link.ContentLength = fi.GetSize() } - return resultLink, nil + return link, nil } func (d *Alias) Other(ctx context.Context, args model.OtherArgs) (interface{}, error) { @@ -511,7 +502,7 @@ func (d *Alias) Extract(ctx context.Context, obj model.Obj, args model.ArchiveIn sign.SignArchive(reqPath)), }, nil } - return link.Clone(), nil + return link, nil } func (d *Alias) ArchiveDecompress(ctx context.Context, srcObj, dstDir model.Obj, args model.ArchiveDecompressArgs) error { diff --git a/drivers/chunk/driver.go b/drivers/chunk/driver.go index 1ebc8aef6c..bf7690887c 100644 --- a/drivers/chunk/driver.go +++ b/drivers/chunk/driver.go @@ -7,6 +7,7 @@ import ( "fmt" "io" stdpath "path" + "slices" "strconv" "strings" @@ -314,7 +315,7 @@ func (d *Chunk) Link(ctx context.Context, file model.Obj, args model.LinkArgs) ( if err != nil { return nil, err } - return l.Clone(), nil + return l, nil } // 检查0号块不等于-1 以支持空文件 // 如果块数量大于1 最后一块不可能为0 @@ -328,7 +329,7 @@ func (d *Chunk) Link(ctx context.Context, file model.Obj, args model.LinkArgs) ( } } fileSize := chunkFile.GetSize() - mergedRrf := func(ctx context.Context, httpRange http_range.Range) (io.ReadCloser, error) { + mergedRrf := func(ctx context.Context, httpRange http_range.Range) (_ io.ReadCloser, err error) { start := httpRange.Start length := httpRange.Length if length < 0 || start+length > fileSize { @@ -339,6 +340,12 @@ func (d *Chunk) Link(ctx context.Context, file model.Obj, args model.LinkArgs) ( } rs := make([]io.Reader, 0) cs := make(utils.Closers, 0) + defer func() { + slices.Reverse(cs) + if err != nil { + err = errors.Join(err, cs.Close()) + } + }() var ( rc io.ReadCloser readFrom bool @@ -347,7 +354,6 @@ func (d *Chunk) Link(ctx context.Context, file model.Obj, args model.LinkArgs) ( if readFrom { l, o, err := op.Link(ctx, remoteStorage, stdpath.Join(remoteActualPath, d.getPartName(idx)), args) if err != nil { - _ = cs.Close() return nil, err } cs = append(cs, l) @@ -356,12 +362,10 @@ func (d *Chunk) Link(ctx context.Context, file model.Obj, args model.LinkArgs) ( chunkSize2 = o.GetSize() } if chunkSize2 != chunkSize { - _ = cs.Close() return nil, fmt.Errorf("chunk part[%d] size not match", idx) } rrf, err := stream.GetRangeReaderFromLink(chunkSize2, l) if err != nil { - _ = cs.Close() return nil, err } newLength := length - chunkSize2 @@ -372,7 +376,9 @@ func (d *Chunk) Link(ctx context.Context, file model.Obj, args model.LinkArgs) ( rc, err = rrf.RangeRead(ctx, http_range.Range{Length: length}) } if err != nil { - _ = cs.Close() + if rc != nil { + err = errors.Join(err, rc.Close()) + } return nil, err } rs = append(rs, rc) @@ -388,7 +394,6 @@ func (d *Chunk) Link(ctx context.Context, file model.Obj, args model.LinkArgs) ( } else { l, o, err := op.Link(ctx, remoteStorage, stdpath.Join(remoteActualPath, d.getPartName(idx)), args) if err != nil { - _ = cs.Close() return nil, err } cs = append(cs, l) @@ -397,17 +402,17 @@ func (d *Chunk) Link(ctx context.Context, file model.Obj, args model.LinkArgs) ( chunkSize2 = o.GetSize() } if chunkSize2 != chunkSize { - _ = cs.Close() return nil, fmt.Errorf("chunk part[%d] size not match", idx) } rrf, err := stream.GetRangeReaderFromLink(chunkSize2, l) if err != nil { - _ = cs.Close() return nil, err } rc, err = rrf.RangeRead(ctx, http_range.Range{Start: start, Length: -1}) if err != nil { - _ = cs.Close() + if rc != nil { + err = errors.Join(err, rc.Close()) + } return nil, err } length -= chunkSize2 - start diff --git a/drivers/crypt/driver.go b/drivers/crypt/driver.go index 7dca84283a..79b5d9a620 100644 --- a/drivers/crypt/driver.go +++ b/drivers/crypt/driver.go @@ -283,10 +283,11 @@ func (d *Crypt) Link(ctx context.Context, file model.Obj, _ model.LinkArgs) (*mo n, err := io.ReadFull(remoteReader, fileHeader) if n != fileHeaderSize { fileHeader = nil + _ = remoteReader.Close() return nil, fmt.Errorf("failed to read all data: (expect =%d, actual =%d) %w", fileHeaderSize, n, err) } if limit <= fileHeaderSize { - remoteReader.Close() + _ = remoteReader.Close() return io.NopCloser(bytes.NewReader(fileHeader[:limit])), nil } else { remoteReader = utils.ReadCloser{ diff --git a/drivers/strm/driver.go b/drivers/strm/driver.go index 3d243c1574..6792518d8b 100644 --- a/drivers/strm/driver.go +++ b/drivers/strm/driver.go @@ -221,7 +221,7 @@ func (d *Strm) Link(ctx context.Context, file model.Obj, args model.LinkArgs) (* }, nil } - return link.Clone(), nil + return link, nil } var _ driver.Driver = (*Strm)(nil) diff --git a/drivers/strm/hook.go b/drivers/strm/hook.go index 24b31ee3e9..5e4d33ed4a 100644 --- a/drivers/strm/hook.go +++ b/drivers/strm/hook.go @@ -124,11 +124,14 @@ func generateStrm(ctx context.Context, driver *Strm, obj model.Obj, localPath st } rc, err := rrf.RangeRead(ctx, http_range.Range{Length: -1}) if err != nil { + if rc != nil { + _ = rc.Close() + } log.Warnf("failed to generate strm of obj %s: failed to read range: %v", localPath, err) return } - defer rc.Close() same, err := isSameContent(localPath, size, rc) + _ = rc.Close() if err != nil { log.Warnf("failed to compare content of obj %s: %v", localPath, err) return @@ -138,6 +141,9 @@ func generateStrm(ctx context.Context, driver *Strm, obj model.Obj, localPath st } rc, err = rrf.RangeRead(ctx, http_range.Range{Length: -1}) if err != nil { + if rc != nil { + _ = rc.Close() + } log.Warnf("failed to generate strm of obj %s: failed to reread range: %v", localPath, err) return } diff --git a/internal/fs/copy_move.go b/internal/fs/copy_move.go index 1d171a9b9c..0cc0673418 100644 --- a/internal/fs/copy_move.go +++ b/internal/fs/copy_move.go @@ -254,7 +254,6 @@ func (t *FileTransferTask) RunWithNextTaskCallback(f func(nextTask *FileTransfer Ctx: t.Ctx(), }, link) if err != nil { - _ = link.Close() return errors.WithMessagef(err, "failed get [%s] stream", t.SrcActualPath) } t.SetTotalBytes(ss.GetSize()) diff --git a/internal/hybrid_cache/buffer.go b/internal/hybrid_cache/buffer.go index 0023996e63..51e6b8834f 100644 --- a/internal/hybrid_cache/buffer.go +++ b/internal/hybrid_cache/buffer.go @@ -3,33 +3,33 @@ package hybrid_cache import ( "fmt" "io" + "sync" ) type BufferStore struct { + mu sync.RWMutex blocks [][]byte size int64 } func (m *BufferStore) Size() int64 { + m.mu.RLock() + defer m.mu.RUnlock() return m.size } -// 用于存储不复用的[]byte -func (m *BufferStore) Append(buf []byte) { - m.size += int64(len(buf)) - m.blocks = append(m.blocks, buf) -} - func (m *BufferStore) Close() error { - if len(m.blocks) > 0 { - clear(m.blocks) - m.blocks = m.blocks[:0] - m.size = 0 - } + m.mu.Lock() + defer m.mu.Unlock() + clear(m.blocks) + m.blocks = nil + m.size = 0 return nil } func (m *BufferStore) ReadAt(p []byte, off int64) (int, error) { + m.mu.RLock() + defer m.mu.RUnlock() if len(p) == 0 { return 0, nil } @@ -55,6 +55,8 @@ func (m *BufferStore) ReadAt(p []byte, off int64) (int, error) { } func (m *BufferStore) WriteAt(p []byte, off int64) (int, error) { + m.mu.RLock() + defer m.mu.RUnlock() if len(p) == 0 { return 0, nil } @@ -80,6 +82,8 @@ func (m *BufferStore) WriteAt(p []byte, off int64) (int, error) { } func (m *BufferStore) GrowTo(size int64) (err error) { + m.mu.Lock() + defer m.mu.Unlock() if size <= m.size { return nil } diff --git a/internal/hybrid_cache/buffer_test.go b/internal/hybrid_cache/buffer_test.go index 439d7cb78e..3180806c11 100644 --- a/internal/hybrid_cache/buffer_test.go +++ b/internal/hybrid_cache/buffer_test.go @@ -14,9 +14,9 @@ func TestBufferStore(t *testing.T) { off int64 } bs := &hybrid_cache.BufferStore{} - bs.Append([]byte("github.com")) - bs.Append([]byte("/OpenList")) - bs.Append([]byte("Team/?")) + initial := []byte("github.com/OpenListTeam/?") + _ = bs.GrowTo(int64(len(initial))) + _, _ = bs.WriteAt(initial, 0) b := []byte("OpenList") off := bs.Size() - 1 _ = bs.GrowTo(off + int64(len(b))) diff --git a/internal/model/args.go b/internal/model/args.go index 16a5c1722f..d1b00ed1dd 100644 --- a/internal/model/args.go +++ b/internal/model/args.go @@ -30,7 +30,7 @@ type Link struct { Header http.Header `json:"header"` // needed header (for url) RangeReader RangeReaderIF `json:"-"` // recommended way if can't use URL - Expiration *time.Duration // local cache expiration; not transferred by Clone + Expiration *time.Duration // local cache expiration //for accelerating request, use multi-thread downloading Concurrency int `json:"concurrency"` @@ -42,12 +42,18 @@ type Link struct { RequireReference bool `json:"-"` } -// Clone transfers ownership of l without inheriting its cache expiration. +// Clone transfers ownership of l while isolating its mutable transport and cache metadata. func (l *Link) Clone() *Link { + var expiration *time.Duration + if l.Expiration != nil { + value := *l.Expiration + expiration = &value + } return &Link{ URL: l.URL, - Header: l.Header, + Header: l.Header.Clone(), RangeReader: l.RangeReader, + Expiration: expiration, Concurrency: l.Concurrency, PartSize: l.PartSize, ContentLength: l.ContentLength, diff --git a/internal/model/args_test.go b/internal/model/args_test.go deleted file mode 100644 index a82ae44b63..0000000000 --- a/internal/model/args_test.go +++ /dev/null @@ -1,25 +0,0 @@ -package model - -import ( - "testing" - "time" -) - -func TestLinkCloneTransfersOwnershipWithoutCachePolicy(t *testing.T) { - ttl := time.Minute - source := &Link{URL: "https://example.test/file", Expiration: &ttl} - - clone := source.Clone() - if clone.URL != source.URL { - t.Fatal("clone did not preserve transport data") - } - if clone.Expiration != nil { - t.Fatal("clone inherited source cache policy") - } - if err := clone.Close(); err != nil { - t.Fatal(err) - } - if !source.Expired() { - t.Fatal("closing clone did not release its source") - } -} diff --git a/internal/net/request.go b/internal/net/request.go index 8c5794e204..9c45294f2d 100644 --- a/internal/net/request.go +++ b/internal/net/request.go @@ -16,12 +16,10 @@ import ( "github.com/OpenListTeam/OpenList/v4/internal/conf" "github.com/OpenListTeam/OpenList/v4/internal/errs" hcache "github.com/OpenListTeam/OpenList/v4/internal/hybrid_cache" - "github.com/OpenListTeam/OpenList/v4/internal/model" "github.com/OpenListTeam/OpenList/v4/pkg/buffer" "github.com/OpenListTeam/OpenList/v4/pkg/utils" "github.com/OpenListTeam/OpenList/v4/pkg/http_range" - "github.com/aws/aws-sdk-go/aws/awsutil" log "github.com/sirupsen/logrus" ) @@ -50,12 +48,16 @@ type Downloader struct { // Concurrency of 1 will download the parts sequentially. Concurrency int - //RequestParam HttpRequestParams - HttpClient HttpRequestFunc + OpenPart OpenPartFunc *ConcurrencyLimit } -type HttpRequestFunc func(ctx context.Context, params *HttpRequestParams) (*http.Response, error) +type PartRequest struct { + Range http_range.Range + First bool +} + +type OpenPartFunc func(ctx context.Context, request PartRequest) (io.ReadCloser, error) func NewDownloader(options ...func(*Downloader)) *Downloader { d := &Downloader{ //允许不设置的选项 @@ -68,18 +70,13 @@ func NewDownloader(options ...func(*Downloader)) *Downloader { return d } -// Download The Downloader makes multi-thread http requests to remote URL, each chunk(except last one) has PartSize, -// cache some data, then return Reader with assembled data -// Supports range, do not support unknown FileSize, and will fail if FileSize is incorrect -// memory usage is at about Concurrency*PartSize, use this wisely -func (d Downloader) Download(ctx context.Context, p *HttpRequestParams) (readCloser io.ReadCloser, err error) { - - var finalP HttpRequestParams - awsutil.Copy(&finalP, p) - if finalP.Range.Length < 0 || finalP.Range.Start+finalP.Range.Length > finalP.Size { - finalP.Range.Length = finalP.Size - finalP.Range.Start +// Download opens range parts concurrently and returns them in order. +// It requires a known size and uses about Concurrency*PartSize bytes of memory. +func (d Downloader) Download(ctx context.Context, size int64, requested http_range.Range) (readCloser io.ReadCloser, err error) { + if requested.Length < 0 || requested.Start+requested.Length > size { + requested.Length = size - requested.Start } - impl := downloader{params: &finalP, cfg: d, ctx: ctx} + impl := downloader{requested: requested, cfg: d, ctx: ctx} // Ensures we don't need nil checks later on // 必需的选项 @@ -92,8 +89,8 @@ func (d Downloader) Download(ctx context.Context, p *HttpRequestParams) (readClo if conf.MinFreeMemory > 0 && impl.cfg.PartSize > int(conf.MaxBlockLimit) { impl.cfg.PartSize = int(conf.MaxBlockLimit) } - if impl.cfg.HttpClient == nil { - impl.cfg.HttpClient = DefaultHttpRequestFunc + if impl.cfg.OpenPart == nil { + return nil, errors.New("missing part opener") } return impl.download() @@ -105,8 +102,8 @@ type downloader struct { cancel context.CancelCauseFunc cfg Downloader - params *HttpRequestParams //http request params - chunkCh chan chunk //chunk chanel + requested http_range.Range + chunkCh chan chunk //chunk chanel //wg sync.WaitGroup mu sync.Mutex @@ -160,8 +157,8 @@ func (d *downloader) download() (io.ReadCloser, error) { } maxPart := 1 - if d.params.Range.Length > int64(d.cfg.PartSize) { - maxPart = int((d.params.Range.Length + int64(d.cfg.PartSize) - 1) / int64(d.cfg.PartSize)) + if d.requested.Length > int64(d.cfg.PartSize) { + maxPart = int((d.requested.Length + int64(d.cfg.PartSize) - 1) / int64(d.cfg.PartSize)) } if maxPart < d.cfg.Concurrency { d.cfg.Concurrency = maxPart @@ -169,13 +166,13 @@ func (d *downloader) download() (io.ReadCloser, error) { log.Debugf("cfgConcurrency:%d", d.cfg.Concurrency) if maxPart == 1 { - resp, err := d.cfg.HttpClient(d.ctx, d.params) + body, err := d.cfg.OpenPart(d.ctx, PartRequest{Range: d.requested, First: true}) if err != nil { d.cfg.ConcurrencyLimit.Release() return nil, err } - closeFunc := resp.Body.Close - resp.Body = utils.NewReadCloser(resp.Body, func() error { + closeFunc := body.Close + body = utils.NewReadCloser(body, func() error { d.mu.Lock() defer d.mu.Unlock() var err error @@ -186,19 +183,19 @@ func (d *downloader) download() (io.ReadCloser, error) { } return err }) - return resp.Body, nil + return body, nil } d.ctx, d.cancel = context.WithCancelCause(d.ctx) // workers d.chunkCh = make(chan chunk, d.cfg.Concurrency) - d.pos = d.params.Range.Start - d.maxPos = d.params.Range.Start + d.params.Range.Length + d.pos = d.requested.Start + d.maxPos = d.requested.Start + d.requested.Length d.concurrency = d.cfg.Concurrency var err error - d.hc, err = hcache.NewHybridCache(uint64(d.cfg.PartSize), uint64(d.params.Range.Length)) + d.hc, err = hcache.NewHybridCache(uint64(d.cfg.PartSize), uint64(d.requested.Length)) if err == nil { d.bufMap = make(map[int]*buffer.PipeBuffer, d.cfg.Concurrency) err = d.sendChunkTask(true) @@ -206,7 +203,7 @@ func (d *downloader) download() (io.ReadCloser, error) { if err != nil { d.cancel(err) d.cfg.ConcurrencyLimit.Release() - _ = d.interrupt() + _ = d.interrupt(false) return nil, err } @@ -252,14 +249,14 @@ func (d *downloader) sendChunkTask(newConcurrency bool) (err error) { switch d.nextChunk { case 0: // 最小分片在前面有助视频播放? - firstSize := d.params.Range.Length % finalSize + firstSize := d.requested.Length % finalSize if firstSize > 0 { minSize := finalSize / 2 // 最小分片太小就调整到一半 finalSize = max(firstSize, minSize) } case 1: - firstSize := d.params.Range.Length % finalSize + firstSize := d.requested.Length % finalSize minSize := finalSize / 2 if firstSize > 0 && firstSize < minSize { finalSize += firstSize - minSize @@ -293,10 +290,10 @@ func (d *downloader) sendChunkTask(newConcurrency bool) (err error) { } // when the final reader Close, we interrupt -func (d *downloader) interrupt() error { +func (d *downloader) interrupt(complete bool) error { err := context.Cause(d.ctx) if err == nil { - if d.written.Load() != d.params.Range.Length { + if !complete && d.written.Load() != d.requested.Length { err = fmt.Errorf("interrupted") } } else if errors.Is(err, context.Canceled) { @@ -309,7 +306,6 @@ func (d *downloader) interrupt() error { for _, buf := range d.bufMap { _ = buf.Close() } - d.bufMap = nil } if d.hc != nil { _ = d.hc.Close() @@ -317,7 +313,6 @@ func (d *downloader) interrupt() error { } if d.maxPos != 0 { d.maxPos = 0 - close(d.chunkCh) if d.concurrency > 0 { d.concurrency = -d.concurrency } @@ -377,11 +372,14 @@ func (d *downloader) downloadPart() { // downloadChunk downloads the chunk func (d *downloader) downloadChunk(ch *chunk) bool { log.Debugf("start chunk_%d, %+v", ch.id, ch) - params := d.getParamsFromChunk(ch) + request := PartRequest{ + Range: http_range.Range{Start: ch.start, Length: ch.size}, + First: ch.id == 0, + } var err error for retry := 0; retry <= d.cfg.PartBodyMaxRetries; retry++ { var n int64 - n, err = d.tryDownloadChunk(params, ch) + n, err = d.tryDownloadChunk(request, ch) if err == nil { d.incrWritten(n) log.Debugf("chunk_%d downloaded", ch.id) @@ -401,11 +399,11 @@ func (d *downloader) downloadChunk(ch *chunk) bool { d.incrWritten(n) ch.start += n ch.size -= n - params.Range.Start = ch.start - params.Range.Length = ch.size + request.Range.Start = ch.start + request.Range.Length = ch.size } - log.Warnf("err chunk_%d, object part download error %s, retrying attempt %d. %v", - ch.id, params.URL, retry, err) + log.Warnf("err chunk_%d, retrying attempt %d. %v", + ch.id, retry, err) } else if err == errInfiniteRetry { retry-- } else if err == errCancelConcurrency { @@ -434,29 +432,18 @@ func (d *downloader) delay(ti time.Duration) bool { var errCancelConcurrency = errors.New("") var errInfiniteRetry = errors.New("") -func (d *downloader) tryDownloadChunk(params *HttpRequestParams, ch *chunk) (int64, error) { - resp, err := d.cfg.HttpClient(d.ctx, params) +func (d *downloader) tryDownloadChunk(request PartRequest, ch *chunk) (int64, error) { + body, err := d.cfg.OpenPart(d.ctx, request) if err != nil { - statusCode, ok := errs.UnwrapOrSelf(err).(HttpStatusCodeError) + _, ok := err.(*errNeedRetry) if !ok { return 0, err } - if statusCode == http.StatusRequestedRangeNotSatisfiable { - return 0, err - } if ch.id == 0 { //第1个任务 有限的重试,超过重试就会结束请求 - switch statusCode { - default: - return 0, err - case http.StatusTooManyRequests: - case http.StatusBadGateway: - case http.StatusServiceUnavailable: - case http.StatusGatewayTimeout: - } if !d.delay(time.Millisecond * time.Duration(rand.Uint32N(300)+200)) { return 0, errCancelConcurrency } - return 0, &errNeedRetry{err} + return 0, err } // 来到这 说明第1个分片下载 连接成功了 @@ -491,16 +478,9 @@ func (d *downloader) tryDownloadChunk(params *HttpRequestParams, ch *chunk) (int return 0, errInfiniteRetry } - defer resp.Body.Close() - //only check file size on the first task - if ch.id == 0 { - err = d.checkTotalBytes(resp) - if err != nil { - return 0, err - } - } + defer body.Close() _ = d.sendChunkTask(true) - n, err := utils.CopyWithBuffer(ch.buf, resp.Body) + n, err := utils.CopyWithBuffer(ch.buf, body) if err != nil { return n, &errNeedRetry{err} @@ -512,16 +492,7 @@ func (d *downloader) tryDownloadChunk(params *HttpRequestParams, ch *chunk) (int return n, nil } -func (d *downloader) getParamsFromChunk(ch *chunk) *HttpRequestParams { - var params HttpRequestParams - awsutil.Copy(¶ms, d.params) - - // Get the getBuf byte range of data - params.Range = http_range.Range{Start: ch.start, Length: ch.size} - return ¶ms -} - -func (d *downloader) checkTotalBytes(resp *http.Response) error { +func checkTotalBytes(resp *http.Response, expected int64) error { var err error totalBytes := int64(-1) contentRange := resp.Header.Get("Content-Range") @@ -550,8 +521,8 @@ func (d *downloader) checkTotalBytes(resp *http.Response) error { } } - if totalBytes != d.params.Size && err == nil { - err = fmt.Errorf("expect file size=%d unmatch remote report size=%d, need refresh cache", d.params.Size, totalBytes) + if totalBytes != expected && err == nil { + err = fmt.Errorf("expect file size=%d unmatch remote report size=%d, need refresh cache", expected, totalBytes) } return err @@ -574,38 +545,35 @@ type chunk struct { newConcurrency bool } -func DefaultHttpRequestFunc(ctx context.Context, params *HttpRequestParams) (*http.Response, error) { - header := http_range.ApplyRangeToHttpHeader(params.Range, params.HeaderRef) - return RequestHttp(ctx, "GET", header, params.URL) -} - -func GetRangeReaderHttpRequestFunc(rangeReader model.RangeReaderIF) HttpRequestFunc { - return func(ctx context.Context, params *HttpRequestParams) (*http.Response, error) { - rc, err := rangeReader.RangeRead(ctx, params.Range) - if err != nil { - return nil, err +func OpenHTTPPart(ctx context.Context, url string, header http.Header, size int64, request PartRequest) (io.ReadCloser, error) { + header = http_range.ApplyRangeToHttpHeader(request.Range, header.Clone()) + resp, err := RequestHttp(ctx, http.MethodGet, header, url) + if err != nil { + status, ok := errs.UnwrapOrSelf(err).(HttpStatusCodeError) + if ok && status != http.StatusRequestedRangeNotSatisfiable && + (!request.First || isRetryableInitialStatus(status)) { + return nil, &errNeedRetry{err} + } + return nil, err + } + if request.First { + if err := checkTotalBytes(resp, size); err != nil { + return nil, errors.Join(err, resp.Body.Close()) } - - return &http.Response{ - StatusCode: http.StatusPartialContent, - Status: http.StatusText(http.StatusPartialContent), - Body: rc, - Header: http.Header{ - "Content-Range": {params.Range.ContentRange(params.Size)}, - }, - ContentLength: params.Range.Length, - }, nil } + return resp.Body, nil } -type HttpRequestParams struct { - URL string - //only want data within this range - Range http_range.Range - HeaderRef http.Header - //total file size - Size int64 +func isRetryableInitialStatus(status HttpStatusCodeError) bool { + switch status { + case http.StatusTooManyRequests, http.StatusBadGateway, + http.StatusServiceUnavailable, http.StatusGatewayTimeout: + return true + default: + return false + } } + type errNeedRetry struct { error } @@ -619,6 +587,7 @@ type multiReadCloser struct { maxPos int curBuf *buffer.PipeBuffer d *downloader + read atomic.Int64 } func (mr *multiReadCloser) Read(p []byte) (n int, err error) { @@ -626,6 +595,7 @@ func (mr *multiReadCloser) Read(p []byte) (n int, err error) { return 0, io.EOF } n, err = mr.curBuf.Read(p) + mr.read.Add(int64(n)) // log.Debugf("read_%d read current buffer, n=%d ,err=%+v", mr.rPos, n, err) if err == io.EOF { log.Debugf("read_%d finished current buffer", mr.pos) @@ -641,5 +611,5 @@ func (mr *multiReadCloser) Read(p []byte) (n int, err error) { } func (mr *multiReadCloser) Close() error { - return mr.d.interrupt() + return mr.d.interrupt(mr.read.Load() == mr.d.requested.Length) } diff --git a/internal/net/request_cancel_test.go b/internal/net/request_cancel_test.go index 83ff187acd..cb52ada428 100644 --- a/internal/net/request_cancel_test.go +++ b/internal/net/request_cancel_test.go @@ -1,9 +1,12 @@ package net import ( + "bytes" "context" "errors" + "io" "net/http" + "sync" "testing" "time" @@ -20,13 +23,13 @@ func TestDownloadCancelledAcquisitionReturnsErrorAndReleasesLimit(t *testing.T) d.Concurrency = 2 d.PartSize = 4 d.ConcurrencyLimit = limit - d.HttpClient = func(ctx context.Context, _ *HttpRequestParams) (*http.Response, error) { + d.OpenPart = func(ctx context.Context, _ PartRequest) (io.ReadCloser, error) { return nil, ctx.Err() } }) ctx, cancel := context.WithCancel(context.Background()) cancel() - reader, err := d.Download(ctx, &HttpRequestParams{Size: 16, Range: http_range.Range{Length: -1}}) + reader, err := d.Download(ctx, 16, http_range.Range{Length: -1}) if reader == nil && err == nil { t.Error("cancelled download returned a nil reader and nil error") } @@ -66,14 +69,14 @@ func TestDownloadSinglePartFailureReleasesLimit(t *testing.T) { d := NewDownloader(func(d *Downloader) { d.PartSize = 32 d.ConcurrencyLimit = limit - d.HttpClient = func(ctx context.Context, _ *HttpRequestParams) (*http.Response, error) { + d.OpenPart = func(ctx context.Context, _ PartRequest) (io.ReadCloser, error) { if err := ctx.Err(); err != nil { return nil, err } return nil, upstreamErr } }) - reader, err := d.Download(tc.ctx, &HttpRequestParams{Size: 16, Range: http_range.Range{Length: -1}}) + reader, err := d.Download(tc.ctx, 16, http_range.Range{Length: -1}) if reader != nil || !errors.Is(err, tc.want) { t.Fatalf("single-part failed download = %v, %v; want nil, %v", reader, err, tc.want) } @@ -86,3 +89,60 @@ func TestDownloadSinglePartFailureReleasesLimit(t *testing.T) { }) } } + +func TestDownloadCancellationDoesNotRaceTaskQueueClose(t *testing.T) { + data := []byte("abcdefgh") + overloadStarted := make(chan struct{}) + releaseOverload := make(chan struct{}) + var overloadOnce sync.Once + limit := &ConcurrencyLimit{Limit: 2} + d := NewDownloader(func(d *Downloader) { + d.Concurrency = 2 + d.PartSize = 4 + d.ConcurrencyLimit = limit + d.OpenPart = func(_ context.Context, request PartRequest) (io.ReadCloser, error) { + if request.Range.Start > 0 { + overloadOnce.Do(func() { close(overloadStarted) }) + <-releaseOverload + return nil, &errNeedRetry{HttpStatusCodeError(http.StatusServiceUnavailable)} + } + end := request.Range.Start + request.Range.Length + return io.NopCloser(bytes.NewReader(data[request.Range.Start:end])), nil + } + }) + ctx, cancel := context.WithCancel(context.Background()) + reader, err := d.Download(ctx, int64(len(data)), http_range.Range{Length: -1}) + if err != nil { + t.Fatal(err) + } + readDone := make(chan struct{}) + go func() { + _, _ = io.ReadAll(reader) + close(readDone) + }() + select { + case <-overloadStarted: + case <-time.After(time.Second): + t.Fatal("later range was not requested") + } + cancel() + if err := reader.Close(); err != nil { + t.Fatalf("close cancelled reader: %v", err) + } + close(releaseOverload) + select { + case <-readDone: + case <-time.After(time.Second): + t.Fatal("cancelled reader remained blocked") + } + for range 100 { + limit.mu.Lock() + remaining := limit.Limit + limit.mu.Unlock() + if remaining == 2 { + return + } + time.Sleep(time.Millisecond) + } + t.Fatal("download workers did not release concurrency slots") +} diff --git a/internal/net/request_test.go b/internal/net/request_test.go index 0fdc56eb33..0d7d96f62b 100644 --- a/internal/net/request_test.go +++ b/internal/net/request_test.go @@ -8,7 +8,6 @@ import ( "context" "fmt" "io" - "net/http" "sync" "testing" "time" @@ -33,7 +32,7 @@ func TestDownloadOrder(t *testing.T) { d := NewDownloader(func(d *Downloader) { d.Concurrency = con d.PartSize = partSize - d.HttpClient = downloader.HttpRequest + d.OpenPart = downloader.OpenPart }) var start, length int64 = 2, 10 @@ -41,11 +40,7 @@ func TestDownloadOrder(t *testing.T) { if length2 == -1 { length2 = int64(len(buff)) - start } - req := &HttpRequestParams{ - Range: http_range.Range{Start: start, Length: length}, - Size: int64(len(buff)), - } - readCloser, err := d.Download(context.Background(), req) + readCloser, err := d.Download(context.Background(), int64(len(buff)), http_range.Range{Start: start, Length: length}) if err != nil { t.Fatalf("expect no error, got %v", err) @@ -76,6 +71,49 @@ func TestDownloadOrder(t *testing.T) { } } +func TestDownloadCloseAfterCompleteReadDoesNotWaitForWorkerReceipt(t *testing.T) { + ctx, cancel := context.WithTimeout(context.Background(), time.Second) + defer cancel() + d := NewDownloader(func(d *Downloader) { + d.Concurrency = 2 + d.PartSize = 4 + d.OpenPart = func(ctx context.Context, request PartRequest) (io.ReadCloser, error) { + return &delayedEOFBody{ctx: ctx, data: make([]byte, request.Range.Length)}, nil + } + }) + reader, err := d.Download(ctx, 8, http_range.Range{Length: 8}) + if err != nil { + t.Fatal(err) + } + data, err := io.ReadAll(reader) + if err != nil { + t.Fatal(err) + } + if len(data) != 8 { + t.Fatalf("read %d bytes, want 8", len(data)) + } + if err := reader.Close(); err != nil { + t.Fatalf("close after complete read: %v", err) + } +} + +type delayedEOFBody struct { + ctx context.Context + data []byte + sent bool +} + +func (r *delayedEOFBody) Read(p []byte) (int, error) { + if !r.sent { + r.sent = true + return copy(p, r.data), nil + } + <-r.ctx.Done() + return 0, r.ctx.Err() +} + +func (*delayedEOFBody) Close() error { return nil } + func TestDownloadInterrupt(t *testing.T) { buff := []byte{1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15} buff = append(buff, buff...) @@ -84,19 +122,15 @@ func TestDownloadInterrupt(t *testing.T) { d := NewDownloader(func(d *Downloader) { d.Concurrency = con d.PartSize = partSize - d.HttpClient = downloader.HttpRequest + d.OpenPart = downloader.OpenPart d.ConcurrencyLimit = &ConcurrencyLimit{ Limit: 5, } }) var start, length int64 = 0, int64(len(buff)) - req := &HttpRequestParams{ - Range: http_range.Range{Start: start, Length: length}, - Size: int64(len(buff)), - } ctx, cancel := context.WithCancel(context.Background()) - readCloser, err := d.Download(ctx, req) + readCloser, err := d.Download(ctx, int64(len(buff)), http_range.Range{Start: start, Length: length}) if err != nil { t.Fatalf("expect no error, got %v", err) @@ -116,28 +150,20 @@ func TestHighConcurrency(t *testing.T) { for i := range len(buff) { buff[i] = byte(i % 256) } - downloader, invocations, _ := newDownloadRangeClient(buff) + downloader, _, _ := newDownloadRangeClient(buff) con, partSize := 64, 100 concurrencyLimit := uint32(32) d := NewDownloader(func(d *Downloader) { d.Concurrency = con d.PartSize = partSize - d.HttpClient = downloader.HttpRequest + d.OpenPart = downloader.OpenPart d.ConcurrencyLimit = &ConcurrencyLimit{ Limit: concurrencyLimit, } }) var start, length int64 = 2, 7 << 10 - length2 := length - if length2 == -1 { - length2 = int64(len(buff)) - start - } - req := &HttpRequestParams{ - Range: http_range.Range{Start: start, Length: length}, - Size: int64(len(buff)), - } - readCloser, err := d.Download(context.Background(), req) + readCloser, err := d.Download(context.Background(), int64(len(buff)), http_range.Range{Start: start, Length: length}) if err != nil { t.Fatalf("expect no error, got %v", err) @@ -146,23 +172,22 @@ func TestHighConcurrency(t *testing.T) { if err != nil { t.Fatalf("expect no error, got %v", err) } - if !bytes.Equal(buff[start:start+length2], resultBuf) { + if !bytes.Equal(buff[start:start+length], resultBuf) { t.Error("expect buffer content matches, but got mismatch") } - chunkSize := int(length+int64(partSize)-1) / partSize - if e, a := chunkSize, *invocations; e != a { - t.Errorf("expect %v API calls, got %v", e, a) - } if err := readCloser.Close(); err != nil { t.Errorf("expect no error on close, got %v", err) } for range 100 { time.Sleep(10 * time.Millisecond) - if d.ConcurrencyLimit.Limit == concurrencyLimit { + d.ConcurrencyLimit.mu.Lock() + remaining := d.ConcurrencyLimit.Limit + d.ConcurrencyLimit.mu.Unlock() + if remaining == concurrencyLimit { return } } - t.Errorf("expect concurrency limit to be %v, got %v", concurrencyLimit, d.ConcurrencyLimit.Limit) + t.Error("download workers did not release concurrency slots") } func init() { @@ -182,16 +207,11 @@ func TestDownloadSingle(t *testing.T) { d := NewDownloader(func(d *Downloader) { d.Concurrency = con d.PartSize = partSize - d.HttpClient = downloader.HttpRequest + d.OpenPart = downloader.OpenPart }) var start, length int64 = 2, 10 - req := &HttpRequestParams{ - Range: http_range.Range{Start: start, Length: length}, - Size: int64(len(buff)), - } - - readCloser, err := d.Download(context.Background(), req) + readCloser, err := d.Download(context.Background(), int64(len(buff)), http_range.Range{Start: start, Length: length}) if err != nil { t.Fatalf("expect no error, got %v", err) @@ -222,7 +242,7 @@ func TestDownloadSingle(t *testing.T) { } type downloadCaptureClient struct { - mockedHttpRequest func(params *HttpRequestParams) (*http.Response, error) + openPart func(request PartRequest) (io.ReadCloser, error) GetObjectInvocations int RetrievedRanges []string @@ -230,36 +250,28 @@ type downloadCaptureClient struct { lock sync.Mutex } -func (c *downloadCaptureClient) HttpRequest(ctx context.Context, params *HttpRequestParams) (*http.Response, error) { +func (c *downloadCaptureClient) OpenPart(_ context.Context, request PartRequest) (io.ReadCloser, error) { c.lock.Lock() defer c.lock.Unlock() c.GetObjectInvocations++ - if params.Range.Length != 0 { - c.RetrievedRanges = append(c.RetrievedRanges, fmt.Sprintf("%d-%d", params.Range.Start, params.Range.Length)) + if request.Range.Length != 0 { + c.RetrievedRanges = append(c.RetrievedRanges, fmt.Sprintf("%d-%d", request.Range.Start, request.Range.Length)) } - return c.mockedHttpRequest(params) + return c.openPart(request) } func newDownloadRangeClient(data []byte) (*downloadCaptureClient, *int, *[]string) { capture := &downloadCaptureClient{} - capture.mockedHttpRequest = func(params *HttpRequestParams) (*http.Response, error) { - start, fin := params.Range.Start, params.Range.Start+params.Range.Length - if params.Range.Length == -1 || fin >= int64(len(data)) { + capture.openPart = func(request PartRequest) (io.ReadCloser, error) { + start, fin := request.Range.Start, request.Range.Start+request.Range.Length + if request.Range.Length == -1 || fin >= int64(len(data)) { fin = int64(len(data)) } - bodyBytes := data[start:fin] - - header := &http.Header{} - header.Set("Content-Range", fmt.Sprintf("bytes %d-%d/%d", start, fin-1, len(data))) - return &http.Response{ - Body: io.NopCloser(bytes.NewReader(bodyBytes)), - Header: *header, - ContentLength: int64(len(bodyBytes)), - }, nil + return io.NopCloser(bytes.NewReader(data[start:fin])), nil } return capture, &capture.GetObjectInvocations, &capture.RetrievedRanges diff --git a/internal/offline_download/tool/transfer.go b/internal/offline_download/tool/transfer.go index 7c5bd164a1..71ce96f662 100644 --- a/internal/offline_download/tool/transfer.go +++ b/internal/offline_download/tool/transfer.go @@ -351,7 +351,6 @@ func transferObjFile(t *TransferTask) error { Ctx: t.Ctx(), }, link) if err != nil { - _ = link.Close() return errors.WithMessagef(err, "failed get [%s] stream", t.SrcActualPath) } t.SetTotalBytes(ss.GetSize()) diff --git a/internal/op/archive.go b/internal/op/archive.go index bb3de11a10..dde458df89 100644 --- a/internal/op/archive.go +++ b/internal/op/archive.go @@ -85,7 +85,6 @@ func GetArchiveToolAndStream(ctx context.Context, storage driver.Driver, path st // Get first part stream ss, err := stream.NewSeekableStream(&stream.FileStream{Ctx: ctx, Obj: obj}, l) if err != nil { - _ = l.Close() return nil, nil, nil, errors.WithMessagef(err, "failed get [%s] stream", path) } ret := []*stream.SeekableStream{ss} @@ -120,7 +119,6 @@ func GetArchiveToolAndStream(ctx context.Context, storage driver.Driver, path st } ss1, e := stream.NewSeekableStream(&stream.FileStream{Ctx: ctx, Obj: o1}, l1) if e != nil { - _ = l1.Close() err = errors.WithMessagef(e, "failed get [%s] stream", p) break } @@ -390,9 +388,24 @@ func ArchiveGet(ctx context.Context, storage driver.Driver, path string, args mo } type objWithLink struct { - link *model.Link - obj model.Obj - policy linkCachePolicy + link *model.Link + obj model.Obj +} + +var errConflictingLinkLifecycle = stderrors.New("invalid link lifecycle: expiration cannot be combined with owned closers or RequireReference") + +func admitLink(link *model.Link, obj model.Obj) (*objWithLink, error) { + if link.Expiration != nil && (link.RequireReference || link.SyncClosers.Length() > 0) { + return nil, stderrors.Join(errConflictingLinkLifecycle, link.Close()) + } + return &objWithLink{link: link, obj: obj}, nil +} + +func (ol *objWithLink) acquire() *model.Link { + if ol.link.Expiration != nil || ol.link.SyncClosers.AcquireReference() || !ol.link.RequireReference { + return ol.link.Clone() + } + return nil } var ( @@ -406,8 +419,8 @@ func DriverExtract(ctx context.Context, storage driver.Driver, path string, args } key := stdpath.Join(Key(storage, path), args.InnerPath) if ol, ok := extractCache.Get(key); ok { - if ol.acquire() { - return ol.link, ol.obj, nil + if link := ol.acquire(); link != nil { + return link, ol.obj, nil } } @@ -416,8 +429,8 @@ func DriverExtract(ctx context.Context, storage driver.Driver, path string, args if err != nil { return nil, errors.Wrapf(err, "failed extract archive") } - if ol.policy.expiration != nil { - extractCache.SetWithTTL(key, ol, *ol.policy.expiration) + if ol.link.Expiration != nil { + extractCache.SetWithTTL(key, ol, *ol.link.Expiration) } else { extractCache.SetWithExpirable(key, ol, &ol.link.SyncClosers) } @@ -429,8 +442,8 @@ func DriverExtract(ctx context.Context, storage driver.Driver, path string, args if err != nil { return nil, nil, err } - if ol.acquire() { - return ol.link, ol.obj, nil + if link := ol.acquire(); link != nil { + return link, ol.obj, nil } } } diff --git a/internal/op/fs.go b/internal/op/fs.go index e6278ca0aa..2ef72dd0bb 100644 --- a/internal/op/fs.go +++ b/internal/op/fs.go @@ -245,8 +245,8 @@ func Link(ctx context.Context, storage driver.Driver, path string, args model.Li } key := Key(storage, path) if ol, exists := Cache.linkCache.GetType(key, typeKey); exists { - if ol.acquire() { - return ol.link, ol.obj, nil + if link := ol.acquire(); link != nil { + return link, ol.obj, nil } } @@ -267,8 +267,8 @@ func Link(ctx context.Context, storage driver.Driver, path string, args model.Li if err != nil { return nil, err } - if ol.policy.expiration != nil { - Cache.linkCache.SetTypeWithTTL(key, typeKey, ol, *ol.policy.expiration) + if ol.link.Expiration != nil { + Cache.linkCache.SetTypeWithTTL(key, typeKey, ol, *ol.link.Expiration) } else { Cache.linkCache.SetTypeWithExpirable(key, typeKey, ol, &link.SyncClosers) } @@ -279,8 +279,8 @@ func Link(ctx context.Context, storage driver.Driver, path string, args model.Li if err != nil { return nil, nil, err } - if ol.acquire() { - return ol.link, ol.obj, nil + if link := ol.acquire(); link != nil { + return link, ol.obj, nil } } } diff --git a/internal/op/link_lifecycle.go b/internal/op/link_lifecycle.go deleted file mode 100644 index 71491c97b0..0000000000 --- a/internal/op/link_lifecycle.go +++ /dev/null @@ -1,34 +0,0 @@ -package op - -import ( - "errors" - "time" - - "github.com/OpenListTeam/OpenList/v4/internal/model" -) - -var errConflictingLinkLifecycle = errors.New("invalid link lifecycle: expiration cannot be combined with owned closers or RequireReference") - -type linkCachePolicy struct { - expiration *time.Duration - requireReference bool -} - -func admitLink(link *model.Link, obj model.Obj) (*objWithLink, error) { - if link.Expiration != nil && (link.RequireReference || link.SyncClosers.Length() > 0) { - return nil, errors.Join(errConflictingLinkLifecycle, link.Close()) - } - return &objWithLink{ - link: link, - obj: obj, - policy: linkCachePolicy{ - expiration: link.Expiration, - requireReference: link.RequireReference, - }, - }, nil -} - -func (ol *objWithLink) acquire() bool { - return ol.policy.expiration != nil || - ol.link.SyncClosers.AcquireReference() || !ol.policy.requireReference -} diff --git a/internal/op/link_lifecycle_test.go b/internal/op/link_lifecycle_test.go index 53bb0baf79..90eee34bb1 100644 --- a/internal/op/link_lifecycle_test.go +++ b/internal/op/link_lifecycle_test.go @@ -2,6 +2,7 @@ package op import ( "context" + "net/http" "strings" "sync/atomic" "testing" @@ -51,19 +52,37 @@ func acquireTestLink(t *testing.T, d *linkLifecycleDriver) *model.Link { } func TestLinkLifecycleModes(t *testing.T) { - t.Run("TTL descriptor remains reusable after close", func(t *testing.T) { + t.Run("TTL borrowers are isolated while the descriptor is reused", func(t *testing.T) { resetLinkLifecycleState(t) ttl := time.Minute d := &linkLifecycleDriver{ Storage: model.Storage{MountPath: "/ttl"}, - links: func() *model.Link { return &model.Link{URL: "https://example.test/file", Expiration: &ttl} }, + links: func() *model.Link { + return &model.Link{ + URL: "https://example.test/file", + Header: http.Header{"X-Link": {"cached"}}, + Expiration: &ttl, + } + }, } first := acquireTestLink(t, d) + if first.Expiration == nil || *first.Expiration != ttl { + t.Fatalf("borrower expiration = %v, want %v", first.Expiration, ttl) + } + *first.Expiration = time.Second + first.URL = "https://borrower.test/file" + first.Header.Set("X-Link", "borrower") _ = first.Close() second := acquireTestLink(t, d) - if second.URL != first.URL || d.calls.Load() != 1 { - t.Fatalf("TTL link was not reused: calls=%d", d.calls.Load()) + if second.URL != "https://example.test/file" || second.Header.Get("X-Link") != "cached" { + t.Fatalf("cached descriptor was mutated: URL=%q Header=%q", second.URL, second.Header.Get("X-Link")) + } + if d.calls.Load() != 1 { + t.Fatalf("driver calls = %d, want 1", d.calls.Load()) + } + if second.Expiration == nil || *second.Expiration != ttl { + t.Fatalf("cached expiration was mutated: %v", second.Expiration) } _ = second.Close() }) diff --git a/internal/stream/stream.go b/internal/stream/stream.go index 09db2ea226..2233c9cc84 100644 --- a/internal/stream/stream.go +++ b/internal/stream/stream.go @@ -231,19 +231,30 @@ type SeekableStream struct { *FileStream // should have one of belows to support rangeRead rangeReader model.RangeReaderIF + link *model.Link } -// NewSeekableStream create a SeekableStream from FileStream and Link -// if FileStream.Reader is not nil, use it directly -// else create RangeReader from Link -func NewSeekableStream(fs *FileStream, link *model.Link) (*SeekableStream, error) { +func (ss *SeekableStream) Close() error { + err := ss.FileStream.Close() + if ss.link != nil { + err = errors.Join(err, ss.link.Close()) + } + return err +} + +// NewSeekableStream adopts link and releases it if construction fails. +func NewSeekableStream(fs *FileStream, link *model.Link) (_ *SeekableStream, err error) { + defer func() { + if err != nil && link != nil { + err = errors.Join(err, link.Close()) + } + }() if len(fs.Mimetype) == 0 { fs.Mimetype = utils.GetMimeType(fs.Obj.GetName()) } if fs.Reader != nil { - fs.Add(link) - return &SeekableStream{FileStream: fs}, nil + return &SeekableStream{FileStream: fs, link: link}, nil } if link != nil { @@ -258,6 +269,9 @@ func NewSeekableStream(fs *FileStream, link *model.Link) (*SeekableStream, error if _, ok := rr.(*model.FileRangeReader); ok { var rc io.ReadCloser rc, err = rr.RangeRead(fs.Ctx, http_range.Range{Length: -1}) + if err != nil && rc != nil { + err = errors.Join(err, rc.Close()) + } if err != nil { return nil, err } @@ -265,8 +279,7 @@ func NewSeekableStream(fs *FileStream, link *model.Link) (*SeekableStream, error fs.Add(rc) } fs.size = size - fs.Add(link) - return &SeekableStream{FileStream: fs, rangeReader: rr}, nil + return &SeekableStream{FileStream: fs, rangeReader: rr, link: link}, nil } return nil, fmt.Errorf("illegal seekableStream") } @@ -276,6 +289,9 @@ func NewSeekableStream(fs *FileStream, link *model.Link) (*SeekableStream, error func (ss *SeekableStream) RangeRead(httpRange http_range.Range) (io.Reader, error) { if ss.GetFile() == nil && ss.rangeReader != nil { rc, err := ss.rangeReader.RangeRead(ss.Ctx, httpRange) + if err != nil && rc != nil { + err = errors.Join(err, rc.Close()) + } if err != nil { return nil, err } @@ -299,6 +315,9 @@ func (ss *SeekableStream) generateReader() error { return fmt.Errorf("illegal seekableStream") } rc, err := ss.rangeReader.RangeRead(ss.Ctx, http_range.Range{Length: -1}) + if err != nil && rc != nil { + err = errors.Join(err, rc.Close()) + } if err != nil { return err } diff --git a/internal/stream/stream_test.go b/internal/stream/stream_test.go index 1d8d002e2d..5c8bf8ec97 100644 --- a/internal/stream/stream_test.go +++ b/internal/stream/stream_test.go @@ -2,9 +2,11 @@ package stream_test import ( "bytes" + "context" "errors" "fmt" "io" + "slices" "testing" "github.com/OpenListTeam/OpenList/v4/internal/conf" @@ -14,6 +16,75 @@ import ( "github.com/OpenListTeam/OpenList/v4/pkg/utils" ) +type observedReadCloser struct { + io.Reader + close func() +} + +func (r *observedReadCloser) Close() error { + r.close() + return nil +} + +func TestNewSeekableStreamOwnsLinkOnConstructionFailure(t *testing.T) { + openErr := errors.New("open failed") + var closed []string + link := &model.Link{ + RangeReader: &model.FileRangeReader{RangeReaderIF: stream.RangeReaderFunc(func(context.Context, http_range.Range) (io.ReadCloser, error) { + return &observedReadCloser{Reader: bytes.NewReader(nil), close: func() { closed = append(closed, "body") }}, openErr + })}, + SyncClosers: utils.NewSyncClosers(utils.CloseFunc(func() error { + closed = append(closed, "link") + return nil + })), + } + + got, err := stream.NewSeekableStream(&stream.FileStream{Ctx: t.Context(), Obj: &model.Object{Name: "file", Size: 4}}, link) + if got != nil || !errors.Is(err, openErr) { + t.Fatalf("NewSeekableStream() = %v, %v; want nil, %v", got, err, openErr) + } + if want := []string{"body", "link"}; !slices.Equal(closed, want) { + t.Fatalf("close order = %v, want %v", closed, want) + } +} + +func TestSeekableStreamOwnsRepeatedRangeBodiesAndLink(t *testing.T) { + data := []byte("abcd") + var closed []string + link := &model.Link{ + RangeReader: stream.RangeReaderFunc(func(_ context.Context, requested http_range.Range) (io.ReadCloser, error) { + end := requested.Start + requested.Length + return &observedReadCloser{Reader: bytes.NewReader(data[requested.Start:end]), close: func() { closed = append(closed, "body") }}, nil + }), + SyncClosers: utils.NewSyncClosers(utils.CloseFunc(func() error { + closed = append(closed, "link") + return nil + })), + } + ss, err := stream.NewSeekableStream(&stream.FileStream{Ctx: t.Context(), Obj: &model.Object{Name: "file", Size: int64(len(data))}}, link) + if err != nil { + t.Fatal(err) + } + for _, requested := range []http_range.Range{{Start: 0, Length: 2}, {Start: 2, Length: 2}} { + reader, err := ss.RangeRead(requested) + if err != nil { + t.Fatal(err) + } + if _, err := io.ReadAll(reader); err != nil { + t.Fatal(err) + } + } + if len(closed) != 0 { + t.Fatalf("premature closes = %v", closed) + } + if err := ss.Close(); err != nil { + t.Fatal(err) + } + if want := []string{"body", "body", "link"}; !slices.Equal(closed, want) { + t.Fatalf("close order = %v, want %v", closed, want) + } +} + func TestRangeRead(t *testing.T) { type args struct { httpRange http_range.Range diff --git a/internal/stream/util.go b/internal/stream/util.go index 2947fcbc2b..7357ca3db2 100644 --- a/internal/stream/util.go +++ b/internal/stream/util.go @@ -34,13 +34,12 @@ func GetRangeReaderFromLink(size int64, link *model.Link) (model.RangeReaderIF, down := net.NewDownloader(func(d *net.Downloader) { d.Concurrency = link.Concurrency d.PartSize = link.PartSize - d.HttpClient = net.GetRangeReaderHttpRequestFunc(link.RangeReader) + d.OpenPart = func(ctx context.Context, request net.PartRequest) (io.ReadCloser, error) { + return link.RangeReader.RangeRead(ctx, request.Range) + } }) rangeReader := func(ctx context.Context, httpRange http_range.Range) (io.ReadCloser, error) { - return down.Download(ctx, &net.HttpRequestParams{ - Range: httpRange, - Size: size, - }) + return down.Download(ctx, size, httpRange) } // RangeReader只能在驱动限速 return RangeReaderFunc(rangeReader), nil @@ -54,30 +53,19 @@ func GetRangeReaderFromLink(size int64, link *model.Link) (model.RangeReaderIF, down := net.NewDownloader(func(d *net.Downloader) { d.Concurrency = link.Concurrency d.PartSize = link.PartSize - d.HttpClient = func(ctx context.Context, params *net.HttpRequestParams) (*http.Response, error) { - if ServerDownloadLimit == nil { - return net.DefaultHttpRequestFunc(ctx, params) - } - resp, err := net.DefaultHttpRequestFunc(ctx, params) - if err == nil && resp.Body != nil { - resp.Body = &RateLimitReader{ - Ctx: ctx, - Reader: resp.Body, - Limiter: ServerDownloadLimit, - } - } - return resp, err - } }) rangeReader := func(ctx context.Context, httpRange http_range.Range) (io.ReadCloser, error) { requestHeader, _ := ctx.Value(conf.RequestHeaderKey).(http.Header) header := net.ProcessHeader(requestHeader, link.Header) - return down.Download(ctx, &net.HttpRequestParams{ - Range: httpRange, - Size: size, - URL: link.URL, - HeaderRef: header, - }) + requestDownloader := *down + requestDownloader.OpenPart = func(ctx context.Context, request net.PartRequest) (io.ReadCloser, error) { + body, err := net.OpenHTTPPart(ctx, link.URL, header, size, request) + if err != nil || ServerDownloadLimit == nil { + return body, err + } + return &RateLimitReader{Ctx: ctx, Reader: body, Limiter: ServerDownloadLimit}, nil + } + return requestDownloader.Download(ctx, size, httpRange) } return RangeReaderFunc(rangeReader), nil } diff --git a/server/common/proxy.go b/server/common/proxy.go index 3522fe9972..17423f5bc3 100644 --- a/server/common/proxy.go +++ b/server/common/proxy.go @@ -17,33 +17,31 @@ import ( "github.com/OpenListTeam/OpenList/v4/pkg/utils" ) -func Proxy(w http.ResponseWriter, r *http.Request, link *model.Link, file model.Obj) error { - // if link.MFile != nil { - // attachHeader(w, file, link) - // http.ServeContent(w, r, file.GetName(), file.ModTime(), link.MFile) - // return nil - // } - - if link.Concurrency > 0 || link.PartSize > 0 { - attachHeader(w, file, link) - size := link.ContentLength - if size <= 0 { - size = file.GetSize() +func Proxy(w http.ResponseWriter, r *http.Request, link *model.Link, file model.Obj, proxyRange bool) error { + defer link.Close() + if proxyRange && link.RangeReader == nil && !strings.HasPrefix(link.URL, GetApiUrl(r.Context())+"/") { + size := file.GetSize() + if link.ContentLength > 0 { + size = link.ContentLength } - rrf, _ := stream.GetRangeReaderFromLink(size, link) - if link.RangeReader == nil { - r = r.WithContext(context.WithValue(r.Context(), conf.RequestHeaderKey, r.Header)) + if rangeReader, err := stream.GetRangeReaderFromLink(size, link); err == nil { + link = &model.Link{RangeReader: rangeReader, ContentLength: size} } - return net.ServeHTTP(w, r, file.GetName(), file.ModTime(), size, rrf) } - - if link.RangeReader != nil { + if link.RangeReader != nil || link.Concurrency > 0 || link.PartSize > 0 { attachHeader(w, file, link) size := link.ContentLength if size <= 0 { size = file.GetSize() } - return net.ServeHTTP(w, r, file.GetName(), file.ModTime(), size, link.RangeReader) + rangeReader, err := stream.GetRangeReaderFromLink(size, link) + if err != nil { + return err + } + if link.RangeReader == nil { + r = r.WithContext(context.WithValue(r.Context(), conf.RequestHeaderKey, r.Header)) + } + return net.ServeHTTP(w, r, file.GetName(), file.ModTime(), size, rangeReader) } //transparent proxy @@ -70,7 +68,6 @@ func Proxy(w http.ResponseWriter, r *http.Request, link *model.Link, file model. func attachHeader(w http.ResponseWriter, file model.Obj, link *model.Link) { fileName := file.GetName() w.Header().Set("Content-Disposition", utils.GenerateContentDisposition(fileName)) - w.Header().Set("Content-Type", utils.GetMimeType(fileName)) size := link.ContentLength if size <= 0 { size = file.GetSize() @@ -97,22 +94,6 @@ func GetEtag(file model.Obj, size int64) string { return fmt.Sprintf(`"%x-%x"`, file.ModTime().Unix(), size) } -func ProxyRange(ctx context.Context, link *model.Link, size int64) *model.Link { - if link.RangeReader == nil && !strings.HasPrefix(link.URL, GetApiUrl(ctx)+"/") { - if link.ContentLength > 0 { - size = link.ContentLength - } - rrf, err := stream.GetRangeReaderFromLink(size, link) - if err == nil { - return &model.Link{ - RangeReader: rrf, - ContentLength: size, - } - } - } - return link -} - type InterceptResponseWriter struct { http.ResponseWriter io.Writer diff --git a/server/common/proxy_cancel_test.go b/server/common/proxy_cancel_test.go index 04f37401eb..46cab97d45 100644 --- a/server/common/proxy_cancel_test.go +++ b/server/common/proxy_cancel_test.go @@ -40,7 +40,7 @@ func TestProxyCancelledPartitionedReaderDoesNotPanic(t *testing.T) { ctx, cancel := context.WithCancel(r.Context()) cancel() w := httptest.NewRecorder() - _ = Proxy(w, r.WithContext(ctx), link, file) + _ = Proxy(w, r.WithContext(ctx), link, file, false) if bytes.Contains(w.Body.Bytes(), []byte("0123456789abcdef")) { t.Errorf("cancelled response contained file contents: %q", w.Body.String()) } diff --git a/server/common/proxy_test.go b/server/common/proxy_test.go index 314763b2f5..1ab921f22f 100644 --- a/server/common/proxy_test.go +++ b/server/common/proxy_test.go @@ -4,7 +4,9 @@ import ( "io" "net/http" "net/http/httptest" + "strings" "testing" + "time" "github.com/OpenListTeam/OpenList/v4/internal/conf" "github.com/OpenListTeam/OpenList/v4/internal/model" @@ -31,11 +33,18 @@ func TestProxyOverridesUpstreamContentDisposition(t *testing.T) { Name: "测试文件.rar", Size: int64(len(content)), } - link := &model.Link{URL: upstream.URL} + closed := 0 + link := &model.Link{ + URL: upstream.URL, + SyncClosers: utils.NewSyncClosers(utils.CloseFunc(func() error { + closed++ + return nil + })), + } recorder := httptest.NewRecorder() request := httptest.NewRequest(http.MethodGet, "/sd/example", nil) - err := Proxy(recorder, request, link, file) + err := Proxy(recorder, request, link, file, false) if err != nil { t.Fatalf("Proxy() error = %v", err) } @@ -51,4 +60,42 @@ func TestProxyOverridesUpstreamContentDisposition(t *testing.T) { if got, want := recorder.Body.String(), content; got != want { t.Errorf("body = %q, want %q", got, want) } + if closed != 1 { + t.Errorf("link close count = %d, want 1", closed) + } +} + +func TestProxyRangePreservesRangeAndClosesLink(t *testing.T) { + previousConfig := conf.Conf + conf.Conf = conf.DefaultConfig("data") + t.Cleanup(func() { conf.Conf = previousConfig }) + + upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + http.ServeContent(w, r, "file.txt", time.Time{}, strings.NewReader("abcdef")) + })) + t.Cleanup(upstream.Close) + + for _, tc := range []struct { + name, method, rangeHeader, body string + status int + }{ + {name: "single range", method: http.MethodGet, rangeHeader: "bytes=1-3", body: "bcd", status: http.StatusPartialContent}, + {name: "head", method: http.MethodHead, body: "", status: http.StatusOK}, + } { + t.Run(tc.name, func(t *testing.T) { + closed := 0 + link := &model.Link{URL: upstream.URL, SyncClosers: utils.NewSyncClosers(utils.CloseFunc(func() error { closed++; return nil }))} + request := httptest.NewRequest(tc.method, "/proxy/file.txt", nil) + if tc.rangeHeader != "" { + request.Header.Set("Range", tc.rangeHeader) + } + recorder := httptest.NewRecorder() + if err := Proxy(recorder, request, link, &model.Object{Name: "file.txt", Size: 6}, true); err != nil { + t.Fatal(err) + } + if recorder.Code != tc.status || recorder.Body.String() != tc.body || closed != 1 { + t.Fatalf("status=%d body=%q closes=%d; want %d, %q, 1", recorder.Code, recorder.Body.String(), closed, tc.status, tc.body) + } + }) + } } diff --git a/server/ftp/fsread.go b/server/ftp/fsread.go index 54a3de8f2c..ca4a100f8f 100644 --- a/server/ftp/fsread.go +++ b/server/ftp/fsread.go @@ -47,7 +47,6 @@ func OpenDownload(ctx context.Context, reqPath string, offset int64) (*FileDownl Ctx: ctx, }, link) if err != nil { - _ = link.Close() return nil, err } reader, err := stream.NewReadAtSeeker(ss, offset) diff --git a/server/handles/down.go b/server/handles/down.go index 50025a0b69..cd19cadd99 100644 --- a/server/handles/down.go +++ b/server/handles/down.go @@ -93,7 +93,6 @@ func redirect(c *gin.Context, link *model.Link) { } func proxy(c *gin.Context, link *model.Link, file model.Obj, proxyRange bool) { - defer link.Close() var err error if link.URL != "" && setting.GetBool(conf.ForwardDirectLinkParams) { query := c.Request.URL.Query() @@ -102,15 +101,13 @@ func proxy(c *gin.Context, link *model.Link, file model.Obj, proxyRange bool) { } link.URL, err = utils.InjectQuery(link.URL, query) if err != nil { + _ = link.Close() common.ErrorPage(c, err, 500) return } } - if proxyRange { - link = common.ProxyRange(c, link, file.GetSize()) - } Writer := &common.WrittenResponseWriter{ResponseWriter: c.Writer} - err = common.Proxy(Writer, c.Request, link, file) + err = common.Proxy(Writer, c.Request, link, file, proxyRange) if err == nil { return } diff --git a/server/webdav/webdav.go b/server/webdav/webdav.go index 06d1431ac3..e598b3ab7b 100644 --- a/server/webdav/webdav.go +++ b/server/webdav/webdav.go @@ -275,12 +275,7 @@ func (h *Handler) handleGetHeadPost(w http.ResponseWriter, r *http.Request) (sta if err != nil { return http.StatusInternalServerError, err } - defer link.Close() - - if storage.GetStorage().ProxyRange { - link = common.ProxyRange(ctx, link, fi.GetSize()) - } - err = common.Proxy(w, r, link, fi) + err = common.Proxy(w, r, link, fi, storage.GetStorage().ProxyRange) if err != nil { if statusCode, ok := errs.UnwrapOrSelf(err).(net.HttpStatusCodeError); ok { return int(statusCode), err