From 1462d63a48482d12c2a30467c5ed688987241cdb Mon Sep 17 00:00:00 2001 From: Nostalgia Date: Thu, 24 Sep 2026 01:47:30 +0800 Subject: [PATCH] fix(net): preserve cancellation errors during partitioned downloads (#3090) --- internal/net/request.go | 12 ++-- internal/net/request_cancel_test.go | 88 +++++++++++++++++++++++++++++ server/common/proxy_cancel_test.go | 49 ++++++++++++++++ 3 files changed, 144 insertions(+), 5 deletions(-) create mode 100644 internal/net/request_cancel_test.go create mode 100644 server/common/proxy_cancel_test.go diff --git a/internal/net/request.go b/internal/net/request.go index 0cfa7942e..8c5794e20 100644 --- a/internal/net/request.go +++ b/internal/net/request.go @@ -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 } } diff --git a/internal/net/request_cancel_test.go b/internal/net/request_cancel_test.go new file mode 100644 index 000000000..83ff187ac --- /dev/null +++ b/internal/net/request_cancel_test.go @@ -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) + } + }) + } +} diff --git a/server/common/proxy_cancel_test.go b/server/common/proxy_cancel_test.go new file mode 100644 index 000000000..04f37401e --- /dev/null +++ b/server/common/proxy_cancel_test.go @@ -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()) + } + }() + } +}