From 902a7d6a0b1c3a4cddde9e76840ff9678a2ddb18 Mon Sep 17 00:00:00 2001 From: nostalume Date: Fri, 11 Sep 2026 23:39:53 +0800 Subject: [PATCH] fix(aliyundrive): limit callback concurrency - Share proxy callback admission by Aliyun user identity and hold permits for complete response-body lifetimes. - Retry only verified callback-capacity rejections while preserving direct redirects and server download limiting. - Map exhausted temporary capacity to S3 SlowDown through the merged OpenListTeam gofakes3 module. - Cover shared limits, lifecycle release, cancellation, retry classification, and the S3 HTTP response. Co-authored-by: Codex <267193182+codex@users.noreply.github.com> # Conflicts: # go.mod # go.sum # server/s3/pager.go --- drivers/aliyundrive_open/callback.go | 295 +++++++++++++++ drivers/aliyundrive_open/callback_test.go | 418 ++++++++++++++++++++++ drivers/aliyundrive_open/driver.go | 22 +- drivers/aliyundrive_open/meta.go | 27 +- go.mod | 2 +- go.sum | 4 +- internal/errs/errors.go | 1 + server/s3/backend.go | 2 + server/s3/errors_test.go | 66 ++++ server/s3/utils.go | 8 + 10 files changed, 825 insertions(+), 20 deletions(-) create mode 100644 drivers/aliyundrive_open/callback.go create mode 100644 drivers/aliyundrive_open/callback_test.go create mode 100644 server/s3/errors_test.go diff --git a/drivers/aliyundrive_open/callback.go b/drivers/aliyundrive_open/callback.go new file mode 100644 index 0000000000..54b0e88b8d --- /dev/null +++ b/drivers/aliyundrive_open/callback.go @@ -0,0 +1,295 @@ +package aliyundrive_open + +import ( + "context" + "fmt" + "io" + "math/rand/v2" + "net/http" + "strings" + "sync" + "time" + + "github.com/OpenListTeam/OpenList/v4/internal/conf" + "github.com/OpenListTeam/OpenList/v4/internal/errs" + anet "github.com/OpenListTeam/OpenList/v4/internal/net" + "github.com/OpenListTeam/OpenList/v4/internal/stream" + "github.com/OpenListTeam/OpenList/v4/pkg/http_range" +) + +const ( + defaultCallbackConcurrency = 1 + callbackAcquireTimeout = time.Second + callbackRequestAttempts = 3 + callbackRetryBaseDelay = 200 * time.Millisecond + callbackErrorBodyLimit = 64 << 10 +) + +var callbackLimiters = struct { + sync.Mutex + byUser map[string]*callbackLimiter +}{byUser: make(map[string]*callbackLimiter)} + +type callbackLimiter struct { + userID string + mu sync.Mutex + active int + nextID uint64 + registrations map[uint64]int + changed chan struct{} +} + +type callbackRegistration struct { + limiter *callbackLimiter + id uint64 + once sync.Once +} + +type callbackPermit struct { + limiter *callbackLimiter + once sync.Once +} + +func normalizeCallbackConcurrency(limit int) int { + if limit <= 0 { + return defaultCallbackConcurrency + } + return limit +} + +func registerCallbackLimiter(userID string, limit int) *callbackRegistration { + callbackLimiters.Lock() + defer callbackLimiters.Unlock() + + limiter := callbackLimiters.byUser[userID] + if limiter == nil { + limiter = &callbackLimiter{ + userID: userID, + registrations: make(map[uint64]int), + changed: make(chan struct{}), + } + callbackLimiters.byUser[userID] = limiter + } + limiter.mu.Lock() + limiter.nextID++ + id := limiter.nextID + limiter.registrations[id] = normalizeCallbackConcurrency(limit) + limiter.signalLocked() + limiter.mu.Unlock() + return &callbackRegistration{limiter: limiter, id: id} +} + +func (r *callbackRegistration) unregister() { + if r == nil || r.limiter == nil { + return + } + r.once.Do(func() { + callbackLimiters.Lock() + defer callbackLimiters.Unlock() + r.limiter.mu.Lock() + delete(r.limiter.registrations, r.id) + r.limiter.signalLocked() + if len(r.limiter.registrations) == 0 && r.limiter.active == 0 { + delete(callbackLimiters.byUser, r.limiter.userID) + } + r.limiter.mu.Unlock() + }) +} + +func (r *callbackRegistration) acquire(ctx context.Context) (*callbackPermit, error) { + if r == nil || r.limiter == nil { + return nil, errs.NewErr(errs.TemporaryCapacity, "callback limiter is unavailable") + } + if err := ctx.Err(); err != nil { + return nil, err + } + waitCtx, cancel := context.WithTimeout(ctx, callbackAcquireTimeout) + defer cancel() + for { + r.limiter.mu.Lock() + if r.limiter.active < r.limiter.limitLocked() { + r.limiter.active++ + r.limiter.mu.Unlock() + return &callbackPermit{limiter: r.limiter}, nil + } + changed := r.limiter.changed + r.limiter.mu.Unlock() + select { + case <-ctx.Done(): + return nil, ctx.Err() + case <-waitCtx.Done(): + if err := ctx.Err(); err != nil { + return nil, err + } + return nil, errs.NewErr(errs.TemporaryCapacity, "timed out waiting for callback admission") + case <-changed: + } + } +} + +func (l *callbackLimiter) limitLocked() int { + limit := 0 + for _, registered := range l.registrations { + if limit == 0 || registered < limit { + limit = registered + } + } + return limit +} + +func (l *callbackLimiter) signalLocked() { + close(l.changed) + l.changed = make(chan struct{}) +} + +func (p *callbackPermit) release() { + if p == nil || p.limiter == nil { + return + } + p.once.Do(func() { + callbackLimiters.Lock() + defer callbackLimiters.Unlock() + p.limiter.mu.Lock() + p.limiter.active-- + p.limiter.signalLocked() + if len(p.limiter.registrations) == 0 && p.limiter.active == 0 { + delete(callbackLimiters.byUser, p.limiter.userID) + } + p.limiter.mu.Unlock() + }) +} + +func (d *AliyundriveOpen) callbackRegistration() *callbackRegistration { + if d.callback != nil { + return d.callback + } + if d.ref != nil { + return d.ref.callbackRegistration() + } + return nil +} + +func (d *AliyundriveOpen) callbackRangeReader(url string, size int64) stream.RangeReaderFunc { + return func(ctx context.Context, requested http_range.Range) (io.ReadCloser, error) { + if requested.Length < 0 || requested.Start+requested.Length > size { + requested.Length = size - requested.Start + } + for attempt := 0; attempt < callbackRequestAttempts; attempt++ { + permit, err := d.callbackRegistration().acquire(ctx) + if err != nil { + return nil, err + } + body, retry, err := openCallbackRange(ctx, url, size, requested) + if !retry && err == nil { + return newCallbackBody(ctx, body, permit.release), nil + } + permit.release() + if !retry { + return nil, err + } + if attempt+1 == callbackRequestAttempts { + return nil, errs.NewErr(errs.TemporaryCapacity, "Aliyun callback concurrency limit rejected %d attempts", callbackRequestAttempts) + } + delay := callbackRetryBaseDelay << attempt + delay += time.Duration(rand.Int64N(int64(delay / 2))) + timer := time.NewTimer(delay) + select { + case <-ctx.Done(): + timer.Stop() + return nil, ctx.Err() + case <-timer.C: + } + } + return nil, errs.NewErr(errs.TemporaryCapacity, "callback attempts exhausted") + } +} + +func openCallbackRange(ctx context.Context, url string, size int64, requested http_range.Range) (io.ReadCloser, bool, error) { + requestHeader, _ := ctx.Value(conf.RequestHeaderKey).(http.Header) + header := anet.ProcessHeader(requestHeader, nil) + header = http_range.ApplyRangeToHttpHeader(requested, header) + req, err := http.NewRequestWithContext(ctx, http.MethodGet, url, nil) + if err != nil { + return nil, false, fmt.Errorf("create Aliyun callback request: %w", err) + } + req.Header = header + response, err := anet.HttpClient().Do(req) + if err != nil { + return nil, false, fmt.Errorf("Aliyun callback request failed: %w", err) + } + if response.StatusCode >= http.StatusBadRequest { + defer response.Body.Close() + body, readErr := io.ReadAll(io.LimitReader(response.Body, callbackErrorBodyLimit)) + if readErr != nil { + return nil, false, fmt.Errorf("read Aliyun callback error response: %w", readErr) + } + if isCallbackCapacityRejection(response.StatusCode, body) { + return nil, true, nil + } + message := strings.ReplaceAll(strings.TrimSpace(string(body)), url, "") + return nil, false, fmt.Errorf("Aliyun callback request failed: %w; response: %s", anet.HttpStatusCodeError(response.StatusCode), message) + } + if requested.Start == 0 && requested.Length == size || response.StatusCode == http.StatusPartialContent || callbackContentRangeStartsAt(response.Header, requested.Start) { + return response.Body, false, nil + } + if response.StatusCode == http.StatusOK { + body, rangeErr := anet.GetRangedHttpReader(response.Body, requested.Start, requested.Length) + if rangeErr != nil { + response.Body.Close() + return nil, false, rangeErr + } + return body, false, nil + } + return response.Body, false, nil +} + +func isCallbackCapacityRejection(status int, body []byte) bool { + return status == http.StatusForbidden && + strings.Contains(string(body), "RequestDeniedByCallback") && + strings.Contains(string(body), "ExceedMaxConcurrency") +} + +func callbackContentRangeStartsAt(header http.Header, offset int64) bool { + start, _, err := http_range.ParseContentRange(header.Get("Content-Range")) + return err == nil && start == offset +} + +type callbackBody struct { + body io.ReadCloser + release func() + once sync.Once + mu sync.Mutex + stop func() bool +} + +func newCallbackBody(ctx context.Context, body io.ReadCloser, release func()) *callbackBody { + b := &callbackBody{body: body, release: release} + stop := context.AfterFunc(ctx, func() { _ = b.Close() }) + b.mu.Lock() + b.stop = stop + b.mu.Unlock() + return b +} + +func (b *callbackBody) Read(p []byte) (int, error) { + n, err := b.body.Read(p) + if err != nil { + _ = b.Close() + } + return n, err +} + +func (b *callbackBody) Close() error { + var err error + b.once.Do(func() { + b.mu.Lock() + stop := b.stop + b.mu.Unlock() + if stop != nil { + stop() + } + err = b.body.Close() + b.release() + }) + return err +} diff --git a/drivers/aliyundrive_open/callback_test.go b/drivers/aliyundrive_open/callback_test.go new file mode 100644 index 0000000000..8bcc3eaeca --- /dev/null +++ b/drivers/aliyundrive_open/callback_test.go @@ -0,0 +1,418 @@ +package aliyundrive_open + +import ( + "context" + "errors" + "fmt" + "io" + "net/http" + "net/http/httptest" + "strings" + "sync" + "sync/atomic" + "testing" + "time" + + "github.com/OpenListTeam/OpenList/v4/drivers/base" + "github.com/OpenListTeam/OpenList/v4/internal/conf" + "github.com/OpenListTeam/OpenList/v4/internal/errs" + "github.com/OpenListTeam/OpenList/v4/internal/model" + anet "github.com/OpenListTeam/OpenList/v4/internal/net" + "github.com/OpenListTeam/OpenList/v4/internal/stream" + "github.com/OpenListTeam/OpenList/v4/pkg/http_range" +) + +func TestProxyLinkLimitsActiveCallbackBodies(t *testing.T) { + oldConf := conf.Conf + conf.Conf = &conf.Config{} + t.Cleanup(func() { conf.Conf = oldConf }) + base.InitClient() + + started := make(chan struct{}, 2) + transport := &callbackTrackingTransport{started: started} + client := anet.HttpClient() + oldTransport := client.Transport + client.Transport = transport + t.Cleanup(func() { client.Transport = oldTransport }) + + var server *httptest.Server + server = httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + switch r.URL.Path { + case "/adrive/v1.0/user/getDriveInfo": + _, _ = fmt.Fprint(w, `{"user_id":"user-1","resource_drive_id":"drive-1"}`) + case "/adrive/v1.0/openFile/getDownloadUrl": + _, _ = fmt.Fprintf(w, `{"url":%q}`, server.URL+"/callback") + default: + http.NotFound(w, r) + } + })) + defer server.Close() + + oldAPIURL := API_URL + API_URL = server.URL + defer func() { API_URL = oldAPIURL }() + + d := &AliyundriveOpen{Addition: Addition{AccessToken: "token"}} + if err := d.Init(t.Context()); err != nil { + t.Fatal(err) + } + defer d.Drop(context.Background()) + if d.CallbackConcurrency != defaultCallbackConcurrency { + t.Fatalf("normalized callback concurrency = %d, want %d", d.CallbackConcurrency, defaultCallbackConcurrency) + } + + link, err := d.Link(t.Context(), &model.Object{ID: "file-1", Name: "file", Size: 1}, model.LinkArgs{}) + if err != nil { + t.Fatal(err) + } + if link.RangeReader == nil { + t.Fatal("proxy link must own callback acquisition through a range reader") + } + if _, ok := link.RangeReader.(stream.RateLimitRangeReaderFunc); !ok { + t.Fatalf("proxy range reader type = %T, want server-rate-limited reader", link.RangeReader) + } + direct, err := d.Link(t.Context(), &model.Object{ID: "file-1", Name: "file", Size: 1}, model.LinkArgs{Redirect: true}) + if err != nil { + t.Fatal(err) + } + if direct.URL == "" || direct.RangeReader != nil { + t.Fatal("redirect link must remain URL-only") + } + + first, err := link.RangeReader.RangeRead(t.Context(), http_range.Range{Length: 1}) + if err != nil { + t.Fatal(err) + } + defer first.Close() + select { + case <-started: + case <-time.After(time.Second): + t.Fatal("first callback request did not start") + } + + type result struct { + body interface{ Close() error } + err error + } + secondResult := make(chan result, 1) + go func() { + body, readErr := link.RangeReader.RangeRead(t.Context(), http_range.Range{Length: 1}) + secondResult <- result{body: body, err: readErr} + }() + + select { + case <-started: + t.Fatal("second callback started before the first body was closed") + case <-time.After(100 * time.Millisecond): + } + if err := first.Close(); err != nil { + t.Fatal(err) + } + + var second result + select { + case second = <-secondResult: + case <-time.After(time.Second): + t.Fatal("second callback did not start after the first body was closed") + } + if second.err != nil { + t.Fatal(second.err) + } + if err := second.body.Close(); err != nil { + t.Fatal(err) + } + if got := transport.peak.Load(); got != 1 { + t.Fatalf("peak active callback bodies = %d, want 1", got) + } +} + +type callbackTrackingTransport struct { + active atomic.Int32 + peak atomic.Int32 + started chan<- struct{} +} + +func (t *callbackTrackingTransport) RoundTrip(request *http.Request) (*http.Response, error) { + current := t.active.Add(1) + for { + old := t.peak.Load() + if current <= old || t.peak.CompareAndSwap(old, current) { + break + } + } + t.started <- struct{}{} + return &http.Response{ + StatusCode: http.StatusPartialContent, + Header: http.Header{"Content-Length": {"1"}, "Content-Range": {"bytes 0-0/1"}}, + Body: &trackedCallbackBody{Reader: strings.NewReader("x"), close: func() { t.active.Add(-1) }}, + ContentLength: 1, + Request: request, + }, nil +} + +type trackedCallbackBody struct { + io.Reader + once sync.Once + close func() +} + +func (b *trackedCallbackBody) Close() error { + b.once.Do(b.close) + return nil +} + +func TestCallbackLimiterUsesMinimumRegisteredLimit(t *testing.T) { + firstRegistration := registerCallbackLimiter(t.Name(), 2) + t.Cleanup(firstRegistration.unregister) + first, err := firstRegistration.acquire(t.Context()) + if err != nil { + t.Fatal(err) + } + second, err := firstRegistration.acquire(t.Context()) + if err != nil { + t.Fatal(err) + } + defer first.release() + defer second.release() + + lowerRegistration := registerCallbackLimiter(t.Name(), 1) + t.Cleanup(lowerRegistration.unregister) + acquired := make(chan *callbackPermit, 1) + go func() { + permit, acquireErr := lowerRegistration.acquire(t.Context()) + if acquireErr == nil { + acquired <- permit + } + }() + + first.release() + select { + case permit := <-acquired: + permit.release() + t.Fatal("lowering the shared limit must wait for all excess bodies to drain") + case <-time.After(100 * time.Millisecond): + } + second.release() + select { + case permit := <-acquired: + permit.release() + case <-time.After(time.Second): + t.Fatal("admission did not resume after active bodies drained below the new limit") + } +} + +func TestCallbackLimiterSeparatesUsers(t *testing.T) { + firstUser := registerCallbackLimiter(t.Name()+"-first", 1) + secondUser := registerCallbackLimiter(t.Name()+"-second", 1) + t.Cleanup(firstUser.unregister) + t.Cleanup(secondUser.unregister) + first, err := firstUser.acquire(t.Context()) + if err != nil { + t.Fatal(err) + } + defer first.release() + second, err := secondUser.acquire(t.Context()) + if err != nil { + t.Fatalf("independent user was blocked: %v", err) + } + second.release() +} + +func TestCallbackLimiterReconfigureWaitsForOldBodies(t *testing.T) { + userID := t.Name() + oldRegistration := registerCallbackLimiter(userID, 2) + first, err := oldRegistration.acquire(t.Context()) + if err != nil { + t.Fatal(err) + } + second, err := oldRegistration.acquire(t.Context()) + if err != nil { + t.Fatal(err) + } + oldRegistration.unregister() + + newRegistration := registerCallbackLimiter(userID, 1) + t.Cleanup(newRegistration.unregister) + acquired := make(chan *callbackPermit, 1) + go func() { + permit, acquireErr := newRegistration.acquire(t.Context()) + if acquireErr == nil { + acquired <- permit + } + }() + first.release() + select { + case permit := <-acquired: + permit.release() + t.Fatal("reconfigured limiter admitted while an old body still occupied the new limit") + case <-time.After(100 * time.Millisecond): + } + second.release() + select { + case permit := <-acquired: + permit.release() + case <-time.After(time.Second): + t.Fatal("reconfigured limiter did not admit after old bodies drained") + } +} + +func TestCallbackLimiterDistinguishesTimeoutAndCancellation(t *testing.T) { + registration := registerCallbackLimiter(t.Name(), 1) + t.Cleanup(registration.unregister) + permit, err := registration.acquire(t.Context()) + if err != nil { + t.Fatal(err) + } + defer permit.release() + + started := time.Now() + _, err = registration.acquire(t.Context()) + if !errors.Is(err, errs.TemporaryCapacity) { + t.Fatalf("admission timeout error = %v, want TemporaryCapacity", err) + } + if time.Since(started) < callbackAcquireTimeout { + t.Fatal("admission timed out before the configured wait elapsed") + } + + ctx, cancel := context.WithCancel(t.Context()) + cancel() + _, err = registration.acquire(ctx) + if !errors.Is(err, context.Canceled) || errors.Is(err, errs.TemporaryCapacity) { + t.Fatalf("canceled admission error = %v, want only context.Canceled", err) + } +} + +func TestCallbackCapacityRejectionRequiresBothExactMarkers(t *testing.T) { + tests := []struct { + name string + body string + want bool + }{ + {name: "both", body: `{"code":"RequestDeniedByCallback","message":"ExceedMaxConcurrency"}`, want: true}, + {name: "code only", body: `{"code":"RequestDeniedByCallback"}`}, + {name: "message only", body: `{"message":"ExceedMaxConcurrency"}`}, + {name: "case differs", body: `{"code":"requestdeniedbycallback","message":"ExceedMaxConcurrency"}`}, + } + for _, test := range tests { + t.Run(test.name, func(t *testing.T) { + if got := isCallbackCapacityRejection(http.StatusForbidden, []byte(test.body)); got != test.want { + t.Fatalf("classification = %v, want %v", got, test.want) + } + }) + } + if isCallbackCapacityRejection(http.StatusTooManyRequests, []byte(`RequestDeniedByCallback ExceedMaxConcurrency`)) { + t.Fatal("non-403 response must not be classified as callback capacity") + } +} + +func TestCallbackRangeRetriesOnlyVerifiedCapacityRejections(t *testing.T) { + var requests atomic.Int32 + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + requests.Add(1) + w.WriteHeader(http.StatusForbidden) + _, _ = io.WriteString(w, `{"code":"RequestDeniedByCallback","message":"ExceedMaxConcurrency"}`) + })) + defer server.Close() + + registration := registerCallbackLimiter(t.Name(), 1) + t.Cleanup(registration.unregister) + d := &AliyundriveOpen{callback: registration} + _, err := d.callbackRangeReader(server.URL+"?token=secret", 1).RangeRead(t.Context(), http_range.Range{Length: 1}) + if !errors.Is(err, errs.TemporaryCapacity) { + t.Fatalf("verified rejection error = %v, want TemporaryCapacity", err) + } + if requests.Load() != callbackRequestAttempts { + t.Fatalf("requests = %d, want %d", requests.Load(), callbackRequestAttempts) + } + if strings.Contains(err.Error(), "secret") { + t.Fatal("capacity error leaked the signed callback URL") + } + permit, acquireErr := registration.acquire(t.Context()) + if acquireErr != nil { + t.Fatalf("capacity retries leaked admission: %v", acquireErr) + } + permit.release() + + requests.Store(0) + permanent := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + requests.Add(1) + w.WriteHeader(http.StatusForbidden) + _, _ = io.WriteString(w, `{"code":"RequestDeniedByCallback","message":"denied"}`) + })) + defer permanent.Close() + _, err = d.callbackRangeReader(permanent.URL+"?token=secret", 1).RangeRead(t.Context(), http_range.Range{Length: 1}) + if errors.Is(err, errs.TemporaryCapacity) { + t.Fatalf("permanent 403 error = %v, must not be TemporaryCapacity", err) + } + if requests.Load() != 1 { + t.Fatalf("permanent 403 requests = %d, want 1", requests.Load()) + } + if strings.Contains(err.Error(), "secret") { + t.Fatal("permanent error leaked the signed callback URL") + } +} + +type countingReadCloser struct { + reader io.Reader + closed atomic.Int32 +} + +func (r *countingReadCloser) Read(p []byte) (int, error) { return r.reader.Read(p) } +func (r *countingReadCloser) Close() error { + r.closed.Add(1) + return nil +} + +func TestCallbackBodyReleasesExactlyOnce(t *testing.T) { + underlying := &countingReadCloser{reader: strings.NewReader("x")} + var released atomic.Int32 + body := newCallbackBody(t.Context(), underlying, func() { released.Add(1) }) + _, _ = io.ReadAll(body) + if err := body.Close(); err != nil { + t.Fatal(err) + } + if err := body.Close(); err != nil { + t.Fatal(err) + } + if underlying.closed.Load() != 1 || released.Load() != 1 { + t.Fatalf("close count = %d, release count = %d; want 1, 1", underlying.closed.Load(), released.Load()) + } +} + +type failingReadCloser struct { + closed atomic.Int32 +} + +func (*failingReadCloser) Read([]byte) (int, error) { return 0, errors.New("read failed") } +func (r *failingReadCloser) Close() error { + r.closed.Add(1) + return nil +} + +func TestCallbackBodyReadFailureReleasesPermit(t *testing.T) { + underlying := &failingReadCloser{} + var released atomic.Int32 + body := newCallbackBody(t.Context(), underlying, func() { released.Add(1) }) + if _, err := body.Read(make([]byte, 1)); err == nil { + t.Fatal("read unexpectedly succeeded") + } + if underlying.closed.Load() != 1 || released.Load() != 1 { + t.Fatalf("close count = %d, release count = %d; want 1, 1", underlying.closed.Load(), released.Load()) + } +} + +func TestCallbackBodyCancellationReleasesPermit(t *testing.T) { + ctx, cancel := context.WithCancel(t.Context()) + underlying := &countingReadCloser{reader: strings.NewReader("x")} + released := make(chan struct{}, 1) + _ = newCallbackBody(ctx, underlying, func() { released <- struct{}{} }) + cancel() + select { + case <-released: + case <-time.After(time.Second): + t.Fatal("context cancellation did not release callback admission") + } + if underlying.closed.Load() != 1 { + t.Fatalf("underlying close count = %d, want 1", underlying.closed.Load()) + } +} diff --git a/drivers/aliyundrive_open/driver.go b/drivers/aliyundrive_open/driver.go index ee93b33030..707eb5b23f 100644 --- a/drivers/aliyundrive_open/driver.go +++ b/drivers/aliyundrive_open/driver.go @@ -11,6 +11,7 @@ import ( "github.com/OpenListTeam/OpenList/v4/internal/driver" "github.com/OpenListTeam/OpenList/v4/internal/errs" "github.com/OpenListTeam/OpenList/v4/internal/model" + "github.com/OpenListTeam/OpenList/v4/internal/stream" "github.com/OpenListTeam/OpenList/v4/pkg/utils" "github.com/go-resty/resty/v2" log "github.com/sirupsen/logrus" @@ -22,8 +23,9 @@ type AliyundriveOpen struct { DriveId string - limiter *limiter - ref *AliyundriveOpen + limiter *limiter + ref *AliyundriveOpen + callback *callbackRegistration } func (d *AliyundriveOpen) Config() driver.Config { @@ -35,6 +37,7 @@ func (d *AliyundriveOpen) GetAddition() driver.Additional { } func (d *AliyundriveOpen) Init(ctx context.Context) error { + d.CallbackConcurrency = normalizeCallbackConcurrency(d.CallbackConcurrency) d.limiter = getLimiterForUser(globalLimiterUserID) // First create a globally shared limiter to limit the initial requests. if d.LIVPDownloadFormat == "" { d.LIVPDownloadFormat = "jpeg" @@ -52,6 +55,7 @@ func (d *AliyundriveOpen) Init(ctx context.Context) error { userid := utils.Json.Get(res, "user_id").ToString() d.limiter.free() d.limiter = getLimiterForUser(userid) // Allocate a corresponding limiter for each user. + d.callback = registerCallbackLimiter(userid, d.CallbackConcurrency) return nil } @@ -65,6 +69,10 @@ func (d *AliyundriveOpen) InitReference(storage driver.Driver) error { } func (d *AliyundriveOpen) Drop(ctx context.Context) error { + if d.callback != nil { + d.callback.unregister() + d.callback = nil + } d.limiter.free() d.limiter = nil d.ref = nil @@ -119,10 +127,16 @@ func (d *AliyundriveOpen) Link(ctx context.Context, file model.Obj, args model.L url = utils.Json.Get(res, "streamsUrl", d.LIVPDownloadFormat).ToString() } exp := time.Minute - return &model.Link{ + link := &model.Link{ URL: url, Expiration: &exp, - }, nil + } + if args.Redirect { + return link, nil + } + link.URL = "" + link.RangeReader = stream.RateLimitRangeReaderFunc(d.callbackRangeReader(url, file.GetSize())) + return link, nil } func (d *AliyundriveOpen) MakeDir(ctx context.Context, parentDir model.Obj, dirName string) (model.Obj, error) { diff --git a/drivers/aliyundrive_open/meta.go b/drivers/aliyundrive_open/meta.go index d76fc2eaab..becfef09b0 100644 --- a/drivers/aliyundrive_open/meta.go +++ b/drivers/aliyundrive_open/meta.go @@ -8,19 +8,20 @@ import ( type Addition struct { DriveType string `json:"drive_type" type:"select" options:"default,resource,backup" default:"resource"` driver.RootID - RefreshToken string `json:"refresh_token" required:"true"` - OrderBy string `json:"order_by" type:"select" options:"name,size,updated_at,created_at"` - OrderDirection string `json:"order_direction" type:"select" options:"ASC,DESC"` - UseOnlineAPI bool `json:"use_online_api" default:"true"` - AlipanType string `json:"alipan_type" required:"true" type:"select" default:"default" options:"default,alipanTV"` - APIAddress string `json:"api_url_address" default:"https://api.oplist.org/alicloud/renewapi"` - ClientID string `json:"client_id" help:"Keep it empty if you don't have one"` - ClientSecret string `json:"client_secret" help:"Keep it empty if you don't have one"` - RemoveWay string `json:"remove_way" required:"true" type:"select" options:"trash,delete"` - RapidUpload bool `json:"rapid_upload" help:"If you enable this option, the file will be uploaded to the server first, so the progress will be incorrect"` - InternalUpload bool `json:"internal_upload" help:"If you are using Aliyun ECS is located in Beijing, you can turn it on to boost the upload speed"` - LIVPDownloadFormat string `json:"livp_download_format" type:"select" options:"jpeg,mov" default:"jpeg"` - AccessToken string + RefreshToken string `json:"refresh_token" required:"true"` + OrderBy string `json:"order_by" type:"select" options:"name,size,updated_at,created_at"` + OrderDirection string `json:"order_direction" type:"select" options:"ASC,DESC"` + UseOnlineAPI bool `json:"use_online_api" default:"true"` + AlipanType string `json:"alipan_type" required:"true" type:"select" default:"default" options:"default,alipanTV"` + APIAddress string `json:"api_url_address" default:"https://api.oplist.org/alicloud/renewapi"` + ClientID string `json:"client_id" help:"Keep it empty if you don't have one"` + ClientSecret string `json:"client_secret" help:"Keep it empty if you don't have one"` + RemoveWay string `json:"remove_way" required:"true" type:"select" options:"trash,delete"` + RapidUpload bool `json:"rapid_upload" help:"If you enable this option, the file will be uploaded to the server first, so the progress will be incorrect"` + InternalUpload bool `json:"internal_upload" help:"If you are using Aliyun ECS is located in Beijing, you can turn it on to boost the upload speed"` + LIVPDownloadFormat string `json:"livp_download_format" type:"select" options:"jpeg,mov" default:"jpeg"` + CallbackConcurrency int `json:"callback_concurrency" type:"number" default:"1" help:"Maximum active proxied downloads shared by Aliyun user"` + AccessToken string } var config = driver.Config{ diff --git a/go.mod b/go.mod index 507a223896..c8b089eb12 100644 --- a/go.mod +++ b/go.mod @@ -10,7 +10,7 @@ require ( github.com/KarpelesLab/reflink v1.0.2 github.com/KirCute/zip v1.0.1 github.com/OpenListTeam/go-cache v0.1.0 - github.com/OpenListTeam/gofakes3 v0.8.1 + github.com/OpenListTeam/gofakes3 v0.8.2-0.20260911142347-cd3c030a83b4 github.com/OpenListTeam/sftpd-openlist v1.0.1 github.com/OpenListTeam/tache v0.2.2 github.com/OpenListTeam/times v0.1.0 diff --git a/go.sum b/go.sum index b53aa8c79f..2f3b75f1a5 100644 --- a/go.sum +++ b/go.sum @@ -51,8 +51,8 @@ github.com/OpenListTeam/115-sdk-go v0.2.6 h1:ehXyStvncvn4qRBuknor3kyGZtUmHc0+stj github.com/OpenListTeam/115-sdk-go v0.2.6/go.mod h1:cfvitk2lwe6036iNi2h+iNxwxWDifKZsSvNtrur5BqU= github.com/OpenListTeam/go-cache v0.1.0 h1:eV2+FCP+rt+E4OCJqLUW7wGccWZNJMV0NNkh+uChbAI= github.com/OpenListTeam/go-cache v0.1.0/go.mod h1:AHWjKhNK3LE4rorVdKyEALDHoeMnP8SjiNyfVlB+Pz4= -github.com/OpenListTeam/gofakes3 v0.8.1 h1:uihJ7Zgb4qIafFcXhcm71BzxCyGRIqBVJYg4YOUa6uY= -github.com/OpenListTeam/gofakes3 v0.8.1/go.mod h1:mS9Ywbo6aId6BrRzeYjOIOpK0QDVnMoKOIb0hpaQZ3U= +github.com/OpenListTeam/gofakes3 v0.8.2-0.20260911142347-cd3c030a83b4 h1:Zy7/qg6aCS0OF/FPIoJh9/d0IgcIxpWRvn79ACm2R/Y= +github.com/OpenListTeam/gofakes3 v0.8.2-0.20260911142347-cd3c030a83b4/go.mod h1:mS9Ywbo6aId6BrRzeYjOIOpK0QDVnMoKOIb0hpaQZ3U= github.com/OpenListTeam/gsync v0.1.0 h1:ywzGybOvA3lW8K1BUjKZ2IUlT2FSlzPO4DOazfYXjcs= github.com/OpenListTeam/gsync v0.1.0/go.mod h1:h/Rvv9aX/6CdW/7B8di3xK3xNV8dUg45Fehrd/ksZ9s= github.com/OpenListTeam/reflink v0.0.0-20260701021214-78760eaeafef h1:67uGHancMF/abMrnkc8abVUWQiG73Wk5d8CKt3RzkFo= diff --git a/internal/errs/errors.go b/internal/errs/errors.go index fdf7f2189a..44f7ff56ee 100644 --- a/internal/errs/errors.go +++ b/internal/errs/errors.go @@ -18,6 +18,7 @@ var ( StorageNotInit = errors.New("storage not init") StreamIncomplete = errors.New("upload/download stream incomplete, possible network issue") StreamPeekFail = errors.New("StreamPeekFail") + TemporaryCapacity = errors.New("temporary capacity unavailable") UnknownArchiveFormat = errors.New("unknown archive format") WrongArchivePassword = errors.New("wrong archive password") diff --git a/server/s3/backend.go b/server/s3/backend.go index be1918d225..ed6c62440d 100644 --- a/server/s3/backend.go +++ b/server/s3/backend.go @@ -152,6 +152,8 @@ func (b *s3Backend) HeadObject(ctx context.Context, bucketName, objectName strin // GetObject fetchs the object from the filesystem. func (b *s3Backend) GetObject(ctx context.Context, bucketName, objectName string, rangeRequest *gofakes3.ObjectRangeRequest) (s3Obj *gofakes3.Object, err error) { + defer func() { err = mapBackendError(err) }() + bucket, err := getBucketByName(bucketName) if err != nil { return nil, err diff --git a/server/s3/errors_test.go b/server/s3/errors_test.go new file mode 100644 index 0000000000..85df7fe813 --- /dev/null +++ b/server/s3/errors_test.go @@ -0,0 +1,66 @@ +package s3 + +import ( + "context" + "encoding/xml" + "errors" + "net/http" + "net/http/httptest" + "testing" + + "github.com/OpenListTeam/OpenList/v4/internal/errs" + "github.com/OpenListTeam/gofakes3" + "github.com/OpenListTeam/gofakes3/s3mem" +) + +func TestMapBackendErrorMapsOnlyTemporaryCapacity(t *testing.T) { + capacity := errs.NewErr(errs.TemporaryCapacity, "callback admission timed out") + if got := mapBackendError(capacity); got != gofakes3.ErrSlowDown { + t.Fatalf("capacity error mapped to %v, want %v", got, gofakes3.ErrSlowDown) + } + + permanent := errors.New("permission denied") + if got := mapBackendError(permanent); got != permanent { + t.Fatalf("permanent error mapped to %v, want original error", got) + } + if got := mapBackendError(nil); got != nil { + t.Fatalf("nil error mapped to %v", got) + } +} + +type capacityBackend struct { + gofakes3.Backend +} + +func (b capacityBackend) GetObject(context.Context, string, string, *gofakes3.ObjectRangeRequest) (*gofakes3.Object, error) { + return nil, mapBackendError(errs.NewErr(errs.TemporaryCapacity, "callback admission timed out")) +} + +func TestTemporaryCapacityProducesS3SlowDownResponse(t *testing.T) { + memory := s3mem.New() + if err := memory.CreateBucket(t.Context(), "bucket"); err != nil { + t.Fatal(err) + } + server := httptest.NewServer(gofakes3.New(capacityBackend{Backend: memory}).Server()) + defer server.Close() + + request, err := http.NewRequestWithContext(t.Context(), http.MethodGet, server.URL+"/bucket/object", nil) + if err != nil { + t.Fatal(err) + } + response, err := server.Client().Do(request) + if err != nil { + t.Fatal(err) + } + defer response.Body.Close() + if response.StatusCode != http.StatusServiceUnavailable { + t.Fatalf("status = %d, want %d", response.StatusCode, http.StatusServiceUnavailable) + } + var result gofakes3.ErrorResult + if err := xml.NewDecoder(response.Body).Decode(&result); err != nil { + t.Fatal(err) + } + if result.Code != gofakes3.ErrSlowDown || result.Message != gofakes3.ErrSlowDown.Message() { + t.Fatalf("S3 error = %#v, want SlowDown with standard message", result) + } +} diff --git a/server/s3/utils.go b/server/s3/utils.go index 0191033d2c..6624a07fcd 100644 --- a/server/s3/utils.go +++ b/server/s3/utils.go @@ -5,6 +5,7 @@ package s3 import ( "context" "encoding/json" + stderrors "errors" "strings" "github.com/OpenListTeam/OpenList/v4/internal/conf" @@ -21,6 +22,13 @@ type Bucket struct { Path string `json:"path"` } +func mapBackendError(err error) error { + if stderrors.Is(err, errs.TemporaryCapacity) { + return gofakes3.ErrSlowDown + } + return err +} + const emptyObjectName = "ThisIsAnEmptyFolderInTheS3Bucket" func getAndParseBuckets() ([]Bucket, error) {