mirror of
https://github.com/OpenListTeam/OpenList.git
synced 2026-10-10 04:53:09 +08:00
fix(net): preserve cancellation errors during partitioned downloads (#3090)
This commit is contained in:
@@ -206,7 +206,8 @@ func (d *downloader) download() (io.ReadCloser, error) {
|
||||
if err != nil {
|
||||
d.cancel(err)
|
||||
d.cfg.ConcurrencyLimit.Release()
|
||||
return nil, d.interrupt()
|
||||
_ = d.interrupt()
|
||||
return nil, err
|
||||
}
|
||||
|
||||
d.mu.Lock()
|
||||
@@ -268,10 +269,6 @@ func (d *downloader) sendChunkTask(newConcurrency bool) (err error) {
|
||||
if err != nil {
|
||||
return err // 分片算法错误或者下载中断
|
||||
}
|
||||
if newConcurrency {
|
||||
go d.downloadPart()
|
||||
d.concurrency--
|
||||
}
|
||||
ch := chunk{
|
||||
start: d.pos,
|
||||
size: finalSize,
|
||||
@@ -286,6 +283,11 @@ func (d *downloader) sendChunkTask(newConcurrency bool) (err error) {
|
||||
case <-d.ctx.Done():
|
||||
return context.Cause(d.ctx)
|
||||
case d.chunkCh <- ch:
|
||||
if newConcurrency {
|
||||
// The worker owns the acquired slot only after its chunk is queued.
|
||||
go d.downloadPart()
|
||||
d.concurrency--
|
||||
}
|
||||
return nil
|
||||
}
|
||||
}
|
||||
|
||||
@@ -0,0 +1,88 @@
|
||||
package net
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"net/http"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/OpenListTeam/OpenList/v4/pkg/http_range"
|
||||
)
|
||||
|
||||
func TestDownloadCancelledAcquisitionReturnsErrorAndReleasesLimit(t *testing.T) {
|
||||
const attempts = 32
|
||||
limits := make([]*ConcurrencyLimit, 0, attempts)
|
||||
for range attempts {
|
||||
limit := &ConcurrencyLimit{Limit: 1}
|
||||
limits = append(limits, limit)
|
||||
d := NewDownloader(func(d *Downloader) {
|
||||
d.Concurrency = 2
|
||||
d.PartSize = 4
|
||||
d.ConcurrencyLimit = limit
|
||||
d.HttpClient = func(ctx context.Context, _ *HttpRequestParams) (*http.Response, 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}})
|
||||
if reader == nil && err == nil {
|
||||
t.Error("cancelled download returned a nil reader and nil error")
|
||||
}
|
||||
if reader != nil {
|
||||
_ = reader.Close()
|
||||
} else if !errors.Is(err, context.Canceled) {
|
||||
t.Errorf("cancelled download error = %v, want context.Canceled", err)
|
||||
}
|
||||
}
|
||||
time.Sleep(50 * time.Millisecond) // allow any started workers to release their slots
|
||||
for i, limit := range limits {
|
||||
limit.mu.Lock()
|
||||
got := limit.Limit
|
||||
limit.mu.Unlock()
|
||||
if got != 1 {
|
||||
t.Errorf("attempt %d remaining concurrency = %d, want 1", i, got)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestDownloadSinglePartFailureReleasesLimit(t *testing.T) {
|
||||
upstreamErr := errors.New("upstream failure")
|
||||
for _, tc := range []struct {
|
||||
name string
|
||||
ctx context.Context
|
||||
want error
|
||||
}{
|
||||
{name: "cancelled", ctx: func() context.Context {
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
cancel()
|
||||
return ctx
|
||||
}(), want: context.Canceled},
|
||||
{name: "upstream failure", ctx: context.Background(), want: upstreamErr},
|
||||
} {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
limit := &ConcurrencyLimit{Limit: 1}
|
||||
d := NewDownloader(func(d *Downloader) {
|
||||
d.PartSize = 32
|
||||
d.ConcurrencyLimit = limit
|
||||
d.HttpClient = func(ctx context.Context, _ *HttpRequestParams) (*http.Response, 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}})
|
||||
if reader != nil || !errors.Is(err, tc.want) {
|
||||
t.Fatalf("single-part failed download = %v, %v; want nil, %v", reader, err, tc.want)
|
||||
}
|
||||
limit.mu.Lock()
|
||||
got := limit.Limit
|
||||
limit.mu.Unlock()
|
||||
if got != 1 {
|
||||
t.Errorf("remaining concurrency = %d, want 1", got)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,49 @@
|
||||
package common
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"io"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"testing"
|
||||
|
||||
"github.com/OpenListTeam/OpenList/v4/internal/conf"
|
||||
"github.com/OpenListTeam/OpenList/v4/internal/model"
|
||||
"github.com/OpenListTeam/OpenList/v4/internal/stream"
|
||||
"github.com/OpenListTeam/OpenList/v4/pkg/http_range"
|
||||
)
|
||||
|
||||
func TestProxyCancelledPartitionedReaderDoesNotPanic(t *testing.T) {
|
||||
oldConf := conf.Conf
|
||||
conf.Conf = conf.DefaultConfig("data")
|
||||
t.Cleanup(func() { conf.Conf = oldConf })
|
||||
link := &model.Link{
|
||||
Concurrency: 2,
|
||||
PartSize: 4,
|
||||
RangeReader: stream.RangeReaderFunc(func(ctx context.Context, requested http_range.Range) (io.ReadCloser, error) {
|
||||
if err := ctx.Err(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return io.NopCloser(bytes.NewReader([]byte("0123456789abcdef")[requested.Start : requested.Start+requested.Length])), nil
|
||||
}),
|
||||
}
|
||||
file := &model.Object{Name: "fixture.bin", Size: 16}
|
||||
for range 32 {
|
||||
func() {
|
||||
defer func() {
|
||||
if recovered := recover(); recovered != nil {
|
||||
t.Errorf("Proxy panicked on cancelled partitioned read: %v", recovered)
|
||||
}
|
||||
}()
|
||||
r := httptest.NewRequest(http.MethodGet, "/proxy/fixture.bin", nil)
|
||||
ctx, cancel := context.WithCancel(r.Context())
|
||||
cancel()
|
||||
w := httptest.NewRecorder()
|
||||
_ = Proxy(w, r.WithContext(ctx), link, file)
|
||||
if bytes.Contains(w.Body.Bytes(), []byte("0123456789abcdef")) {
|
||||
t.Errorf("cancelled response contained file contents: %q", w.Body.String())
|
||||
}
|
||||
}()
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user