From 893457cd50d954b4a1a81727273b073193b1997d Mon Sep 17 00:00:00 2001 From: Nostalgia Date: Thu, 24 Sep 2026 12:01:50 +0800 Subject: [PATCH] fix(aliyundrive): limit callback concurrency (#3071) * 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 * fix(op): separate redirect and proxy link cache entries - Include redirect mode in the link cache key for all drivers. - Cover both redirect-to-proxy and proxy-to-redirect cache reuse. Co-authored-by: Codex <267193182+codex@users.noreply.github.com> * fix(proxy): close range bodies before opening next - make ServeHTTP own each range body and preserve cleanup failures - remove the aggregate range closer and pass range readers directly - replace the obsolete callback transport test with focused lifecycle coverage Co-authored-by: Codex <267193182+codex@users.noreply.github.com> --------- Co-authored-by: nostalume Co-authored-by: Codex <267193182+codex@users.noreply.github.com> --- drivers/aliyundrive_open/callback.go | 295 +++++++++++++++++ drivers/aliyundrive_open/callback_test.go | 365 ++++++++++++++++++++++ drivers/aliyundrive_open/driver.go | 22 +- drivers/aliyundrive_open/meta.go | 27 +- go.mod | 2 +- go.sum | 4 +- internal/errs/errors.go | 1 + internal/model/args.go | 18 -- internal/net/serve.go | 78 +++-- internal/net/serve_test.go | 236 ++++++++++++++ internal/op/fs.go | 5 +- internal/op/link_cache_test.go | 71 +++++ server/common/proxy.go | 8 +- server/s3/backend.go | 2 + server/s3/errors_test.go | 66 ++++ server/s3/utils.go | 8 + 16 files changed, 1139 insertions(+), 69 deletions(-) create mode 100644 drivers/aliyundrive_open/callback.go create mode 100644 drivers/aliyundrive_open/callback_test.go create mode 100644 internal/net/serve_test.go create mode 100644 internal/op/link_cache_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 000000000..54b0e88b8 --- /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 000000000..d16d22cc8 --- /dev/null +++ b/drivers/aliyundrive_open/callback_test.go @@ -0,0 +1,365 @@ +package aliyundrive_open + +import ( + "context" + "errors" + "fmt" + "io" + "net/http" + "net/http/httptest" + "strings" + "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" + "github.com/OpenListTeam/OpenList/v4/internal/stream" + "github.com/OpenListTeam/OpenList/v4/pkg/http_range" +) + +func TestLinkSeparatesRedirectAndProxyRepresentations(t *testing.T) { + oldConf := conf.Conf + conf.Conf = &conf.Config{} + t.Cleanup(func() { conf.Conf = oldConf }) + base.InitClient() + + 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") + } +} + +func TestCallbackRangeHoldsPermitUntilBodyClose(t *testing.T) { + oldConf := conf.Conf + conf.Conf = &conf.Config{} + t.Cleanup(func() { conf.Conf = oldConf }) + + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + w.Header().Set("Content-Length", "1") + w.Header().Set("Content-Range", "bytes 0-0/1") + w.WriteHeader(http.StatusPartialContent) + _, _ = io.WriteString(w, "x") + })) + defer server.Close() + + registration := registerCallbackLimiter(t.Name(), 1) + t.Cleanup(registration.unregister) + d := &AliyundriveOpen{callback: registration} + body, err := d.callbackRangeReader(server.URL, 1).RangeRead(t.Context(), http_range.Range{Length: 1}) + if err != nil { + t.Fatal(err) + } + registration.limiter.mu.Lock() + active := registration.limiter.active + registration.limiter.mu.Unlock() + if active != 1 { + t.Fatalf("active callback bodies = %d, want 1", active) + } + if err := body.Close(); err != nil { + t.Fatal(err) + } + registration.limiter.mu.Lock() + active = registration.limiter.active + registration.limiter.mu.Unlock() + if active != 0 { + t.Fatalf("active callback bodies after Close = %d, want 0", active) + } +} + +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 ee93b3303..707eb5b23 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 d76fc2eaa..becfef09b 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 507a22389..c8b089eb1 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 b53aa8c79..2f3b75f1a 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 fdf7f2189..44f7ff56e 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/internal/model/args.go b/internal/model/args.go index b051106cc..16a5c1722 100644 --- a/internal/model/args.go +++ b/internal/model/args.go @@ -118,21 +118,3 @@ type SharingLinkArgs struct { type RangeReaderIF interface { RangeRead(ctx context.Context, httpRange http_range.Range) (io.ReadCloser, error) } - -type RangeReadCloserIF interface { - RangeReaderIF - utils.ClosersIF -} - -var _ RangeReadCloserIF = (*RangeReadCloser)(nil) - -type RangeReadCloser struct { - RangeReader RangeReaderIF - utils.Closers -} - -func (r *RangeReadCloser) RangeRead(ctx context.Context, httpRange http_range.Range) (io.ReadCloser, error) { - rc, err := r.RangeReader.RangeRead(ctx, httpRange) - r.Add(rc) - return rc, err -} diff --git a/internal/net/serve.go b/internal/net/serve.go index 89a209d88..5af982bdf 100644 --- a/internal/net/serve.go +++ b/internal/net/serve.go @@ -4,6 +4,7 @@ import ( "compress/gzip" "context" "crypto/tls" + stderrors "errors" "fmt" "io" "mime/multipart" @@ -15,7 +16,6 @@ import ( "time" "github.com/OpenListTeam/OpenList/v4/internal/conf" - "github.com/OpenListTeam/OpenList/v4/internal/errs" "github.com/OpenListTeam/OpenList/v4/internal/model" "github.com/OpenListTeam/OpenList/v4/pkg/http_range" "github.com/OpenListTeam/OpenList/v4/pkg/utils" @@ -25,12 +25,8 @@ import ( //this file is inspired by GO_SDK net.http.ServeContent -//type RangeReadCloser struct { -// GetReaderForRange RangeReaderFunc -//} - // ServeHTTP replies to the request using the content in the -// provided RangeReadCloser. The main benefit of ServeHTTP over io.Copy +// provided range reader. The main benefit of ServeHTTP over io.Copy // is that it handles Range requests properly, sets the MIME type, and // handles If-Match, If-Unmodified-Since, If-None-Match, If-Modified-Since, // and If-Range requests. @@ -47,13 +43,11 @@ import ( // request includes an If-Modified-Since header, ServeHTTP uses // modtime to decide whether the content needs to be sent at all. // -// The content's RangeReadCloser method must work: ServeHTTP gives a range, -// caller will give the reader for that Range. +// The content's RangeRead method must return a reader for the requested range. // // If the caller has set w's ETag header formatted per RFC 7232, section 2.3, // ServeHTTP uses it to handle requests using If-Match, If-None-Match, or If-Range. -func ServeHTTP(w http.ResponseWriter, r *http.Request, name string, modTime time.Time, size int64, RangeReadCloser model.RangeReadCloserIF) error { - defer RangeReadCloser.Close() +func ServeHTTP(w http.ResponseWriter, r *http.Request, name string, modTime time.Time, size int64, rangeReader model.RangeReaderIF) (err error) { setLastModified(w, modTime) done, rangeReq := checkPreconditions(w, r, modTime) if done { @@ -113,10 +107,11 @@ func ServeHTTP(w http.ResponseWriter, r *http.Request, name string, modTime time ctx := r.Context() switch { case len(ranges) == 0: - reader, err := RangeReadCloser.RangeRead(ctx, http_range.Range{Length: -1}) + reader, err := openRange(ctx, rangeReader, http_range.Range{Length: -1}) if err != nil { code = http.StatusRequestedRangeNotSatisfiable - if statusCode, ok := errs.UnwrapOrSelf(err).(HttpStatusCodeError); ok { + var statusCode HttpStatusCodeError + if errors.As(err, &statusCode) { code = int(statusCode) } http.Error(w, err.Error(), code) @@ -136,10 +131,11 @@ func ServeHTTP(w http.ResponseWriter, r *http.Request, name string, modTime time // does not request multiple parts might not support // multipart responses." ra := ranges[0] - sendContent, err = RangeReadCloser.RangeRead(ctx, ra) + sendContent, err = openRange(ctx, rangeReader, ra) if err != nil { code = http.StatusRequestedRangeNotSatisfiable - if statusCode, ok := errs.UnwrapOrSelf(err).(HttpStatusCodeError); ok { + var statusCode HttpStatusCodeError + if errors.As(err, &statusCode) { code = int(statusCode) } http.Error(w, err.Error(), code) @@ -159,7 +155,6 @@ func ServeHTTP(w http.ResponseWriter, r *http.Request, name string, modTime time mw := multipart.NewWriter(pw) w.Header().Set("Content-Type", "multipart/byteranges; boundary="+mw.Boundary()) sendContent = pr - defer pr.Close() // cause writing goroutine to fail and exit if CopyN doesn't finish. go func() { for _, ra := range ranges { part, err := mw.CreatePart(ra.MimeHeader(contentType, size)) @@ -167,21 +162,18 @@ func ServeHTTP(w http.ResponseWriter, r *http.Request, name string, modTime time pw.CloseWithError(err) return } - reader, err := RangeReadCloser.RangeRead(ctx, ra) - if err != nil { - pw.CloseWithError(err) - return - } - if _, err := utils.CopyWithBufferN(part, reader, ra.Length); err != nil { + if err := copyRange(ctx, part, rangeReader, ra); err != nil { pw.CloseWithError(err) return } } - mw.Close() - pw.Close() + _ = pw.CloseWithError(mw.Close()) }() } + defer func() { + err = closeWithError(err, sendContent) + }() w.Header().Set("Accept-Ranges", "bytes") if w.Header().Get("Content-Encoding") == "" { @@ -201,7 +193,8 @@ func ServeHTTP(w http.ResponseWriter, r *http.Request, name string, modTime time log.Warnf("Maybe size incorrect or reader not giving correct/full data, or connection closed before finish. written bytes: %d ,sendSize:%d, ", written, sendSize) } code = http.StatusInternalServerError - if statusCode, ok := errs.UnwrapOrSelf(err).(HttpStatusCodeError); ok { + var statusCode HttpStatusCodeError + if errors.As(err, &statusCode) { code = int(statusCode) } w.WriteHeader(code) @@ -210,6 +203,43 @@ func ServeHTTP(w http.ResponseWriter, r *http.Request, name string, modTime time } return nil } + +func copyRange(ctx context.Context, dst io.Writer, rangeReader model.RangeReaderIF, requested http_range.Range) (err error) { + reader, err := openRange(ctx, rangeReader, requested) + if err != nil { + return err + } + defer func() { + err = closeWithError(err, reader) + }() + _, err = utils.CopyWithBufferN(dst, reader, requested.Length) + return err +} + +func openRange(ctx context.Context, rangeReader model.RangeReaderIF, requested http_range.Range) (io.ReadCloser, error) { + reader, err := rangeReader.RangeRead(ctx, requested) + if err != nil { + if reader != nil { + err = closeWithError(err, reader) + } + return nil, err + } + if reader == nil { + return nil, errors.New("range reader returned a nil body") + } + return reader, nil +} + +func closeWithError(err error, closer io.Closer) error { + closeErr := closer.Close() + if err == nil { + return closeErr + } + if closeErr == nil { + return err + } + return stderrors.Join(err, closeErr) +} func ProcessHeader(origin, override http.Header) http.Header { result := http.Header{} // client header diff --git a/internal/net/serve_test.go b/internal/net/serve_test.go new file mode 100644 index 000000000..d28d2ed1d --- /dev/null +++ b/internal/net/serve_test.go @@ -0,0 +1,236 @@ +package net + +import ( + "bytes" + "context" + "errors" + "fmt" + "io" + "mime" + "mime/multipart" + "net/http" + "net/http/httptest" + "reflect" + "sync" + "testing" + "time" + + "github.com/OpenListTeam/OpenList/v4/pkg/http_range" +) + +func TestServeHTTPClosesMultipartRangeBeforeOpeningNext(t *testing.T) { + source := newSequentialRangeSource("abc") + ctx, cancel := context.WithTimeout(t.Context(), 500*time.Millisecond) + defer cancel() + request := httptest.NewRequest(http.MethodGet, "/file", nil).WithContext(ctx) + request.Header.Set("Range", "bytes=0-0,2-2") + recorder := httptest.NewRecorder() + + if err := ServeHTTP(recorder, request, "file.txt", time.Time{}, 3, source); err != nil { + t.Fatalf("ServeHTTP() error = %v", err) + } + response := recorder.Result() + defer response.Body.Close() + if response.StatusCode != http.StatusPartialContent { + t.Fatalf("status = %d, want %d", response.StatusCode, http.StatusPartialContent) + } + + mediaType, params, err := mime.ParseMediaType(response.Header.Get("Content-Type")) + if err != nil { + t.Fatalf("parse Content-Type: %v", err) + } + if mediaType != "multipart/byteranges" { + t.Fatalf("Content-Type = %q, want multipart/byteranges", mediaType) + } + multipartReader := multipart.NewReader(response.Body, params["boundary"]) + var parts []string + for { + part, err := multipartReader.NextPart() + if errors.Is(err, io.EOF) { + break + } + if err != nil { + t.Fatalf("read multipart part: %v", err) + } + body, err := io.ReadAll(part) + if err != nil { + t.Fatalf("read multipart body: %v", err) + } + parts = append(parts, string(body)) + } + if want := []string{"a", "c"}; !reflect.DeepEqual(parts, want) { + t.Fatalf("multipart parts = %q, want %q", parts, want) + } + assertRangeLifecycle(t, source, []string{"open:0", "close:0", "open:2", "close:2"}, []int{1, 1}) +} + +func TestServeHTTPClosesSelectedRangeBody(t *testing.T) { + tests := []struct { + name string + method string + rangeValue string + wantStatus int + wantEvents []string + }{ + {name: "full", method: http.MethodGet, wantStatus: http.StatusOK, wantEvents: []string{"open:0", "close:0"}}, + {name: "single range", method: http.MethodGet, rangeValue: "bytes=1-1", wantStatus: http.StatusPartialContent, wantEvents: []string{"open:1", "close:1"}}, + {name: "head", method: http.MethodHead, wantStatus: http.StatusOK, wantEvents: []string{"open:0", "close:0"}}, + } + for _, test := range tests { + t.Run(test.name, func(t *testing.T) { + source := newSequentialRangeSource("abc") + request := httptest.NewRequest(test.method, "/file", nil) + if test.rangeValue != "" { + request.Header.Set("Range", test.rangeValue) + } + recorder := httptest.NewRecorder() + + if err := ServeHTTP(recorder, request, "file.txt", time.Time{}, 3, source); err != nil { + t.Fatalf("ServeHTTP() error = %v", err) + } + if recorder.Code != test.wantStatus { + t.Fatalf("status = %d, want %d", recorder.Code, test.wantStatus) + } + assertRangeLifecycle(t, source, test.wantEvents, []int{1}) + }) + } +} + +func TestServeHTTPClosesRangeAfterWriteFailure(t *testing.T) { + writeErr := errors.New("write failed") + source := newSequentialRangeSource("abc") + request := httptest.NewRequest(http.MethodGet, "/file", nil) + writer := &failingResponseWriter{header: make(http.Header), err: writeErr} + + err := ServeHTTP(writer, request, "file.txt", time.Time{}, 3, source) + if !errors.Is(err, writeErr) { + t.Fatalf("ServeHTTP() error = %v, want %v", err, writeErr) + } + assertRangeLifecycle(t, source, []string{"open:0", "close:0"}, []int{1}) +} + +func TestServeHTTPClosesBodyReturnedWithOpenError(t *testing.T) { + source := newSequentialRangeSource("abc") + source.openErr = HttpStatusCodeError(http.StatusServiceUnavailable) + source.closeErr = errors.New("close failed") + request := httptest.NewRequest(http.MethodGet, "/file", nil) + recorder := httptest.NewRecorder() + + if err := ServeHTTP(recorder, request, "file.txt", time.Time{}, 3, source); err != nil { + t.Fatalf("ServeHTTP() error = %v", err) + } + if recorder.Code != http.StatusServiceUnavailable { + t.Fatalf("status = %d, want %d", recorder.Code, http.StatusServiceUnavailable) + } + assertRangeLifecycle(t, source, []string{"open:0", "close:0"}, []int{1}) +} + +func TestServeHTTPStopsMultipartAfterRangeCloseFailure(t *testing.T) { + closeErr := errors.New("close failed") + source := newSequentialRangeSource("abc") + source.closeErr = closeErr + request := httptest.NewRequest(http.MethodGet, "/file", nil) + request.Header.Set("Range", "bytes=0-0,2-2") + recorder := httptest.NewRecorder() + + err := ServeHTTP(recorder, request, "file.txt", time.Time{}, 3, source) + if !errors.Is(err, closeErr) { + t.Fatalf("ServeHTTP() error = %v, want %v", err, closeErr) + } + assertRangeLifecycle(t, source, []string{"open:0", "close:0"}, []int{1}) +} + +type sequentialRangeSource struct { + content []byte + permit chan struct{} + + mu sync.Mutex + events []string + closeCounts []int + closeErr error + openErr error +} + +func newSequentialRangeSource(content string) *sequentialRangeSource { + return &sequentialRangeSource{ + content: []byte(content), + permit: make(chan struct{}, 1), + } +} + +func (s *sequentialRangeSource) RangeRead(ctx context.Context, requested http_range.Range) (io.ReadCloser, error) { + select { + case s.permit <- struct{}{}: + case <-ctx.Done(): + return nil, ctx.Err() + } + + start := int(requested.Start) + length := int(requested.Length) + if length < 0 || start+length > len(s.content) { + length = len(s.content) - start + } + end := start + length + s.mu.Lock() + index := len(s.closeCounts) + s.events = append(s.events, fmt.Sprintf("open:%d", requested.Start)) + s.closeCounts = append(s.closeCounts, 0) + s.mu.Unlock() + return &testReadCloser{ + Reader: bytes.NewReader(s.content[start:end]), + close: func() error { + s.mu.Lock() + s.closeCounts[index]++ + closeCalls := s.closeCounts[index] + if closeCalls == 1 { + s.events = append(s.events, fmt.Sprintf("close:%d", requested.Start)) + } + s.mu.Unlock() + if closeCalls != 1 { + return fmt.Errorf("body closed %d times", closeCalls) + } + <-s.permit + return s.closeErr + }, + }, s.openErr +} + +func (s *sequentialRangeSource) eventsSnapshot() []string { + s.mu.Lock() + defer s.mu.Unlock() + return append([]string(nil), s.events...) +} + +func (s *sequentialRangeSource) closeCountsSnapshot() []int { + s.mu.Lock() + defer s.mu.Unlock() + return append([]int(nil), s.closeCounts...) +} + +func assertRangeLifecycle(t *testing.T, source *sequentialRangeSource, wantEvents []string, wantCloseCounts []int) { + t.Helper() + if got := source.eventsSnapshot(); !reflect.DeepEqual(got, wantEvents) { + t.Fatalf("range lifecycle = %v, want %v", got, wantEvents) + } + if got := source.closeCountsSnapshot(); !reflect.DeepEqual(got, wantCloseCounts) { + t.Fatalf("close counts = %v, want %v", got, wantCloseCounts) + } +} + +type failingResponseWriter struct { + header http.Header + err error +} + +func (w *failingResponseWriter) Header() http.Header { return w.header } +func (*failingResponseWriter) WriteHeader(int) {} +func (w *failingResponseWriter) Write([]byte) (int, error) { + return 0, w.err +} + +type testReadCloser struct { + io.Reader + close func() error +} + +func (b *testReadCloser) Close() error { return b.close() } diff --git a/internal/op/fs.go b/internal/op/fs.go index 033d89b96..e6278ca0a 100644 --- a/internal/op/fs.go +++ b/internal/op/fs.go @@ -233,7 +233,10 @@ func Link(ctx context.Context, storage driver.Driver, path string, args model.Li if mode == -1 { mode = storage.(driver.LinkCacheModeResolver).ResolveLinkCacheMode(path) } - typeKey := args.Type + typeKey := "proxy/" + args.Type + if args.Redirect { + typeKey = "redirect/" + args.Type + } if mode&driver.LinkCacheIP != 0 { typeKey += "/" + args.IP } diff --git a/internal/op/link_cache_test.go b/internal/op/link_cache_test.go new file mode 100644 index 000000000..6924bf837 --- /dev/null +++ b/internal/op/link_cache_test.go @@ -0,0 +1,71 @@ +package op + +import ( + "context" + "io" + "strings" + "testing" + "time" + + "github.com/OpenListTeam/OpenList/v4/internal/driver" + "github.com/OpenListTeam/OpenList/v4/internal/model" + "github.com/OpenListTeam/OpenList/v4/internal/stream" + "github.com/OpenListTeam/OpenList/v4/pkg/http_range" +) + +type linkModeDriver struct { + driver.Driver + storage model.Storage + calls int +} + +func (d *linkModeDriver) Config() driver.Config { return driver.Config{} } + +func (d *linkModeDriver) GetStorage() *model.Storage { return &d.storage } + +func (d *linkModeDriver) Get(context.Context, string) (model.Obj, error) { + return &model.Object{Name: "file"}, nil +} + +func (d *linkModeDriver) Link(_ context.Context, _ model.Obj, args model.LinkArgs) (*model.Link, error) { + d.calls++ + expiration := time.Minute + if args.Redirect { + return &model.Link{URL: "https://example.com/file", Expiration: &expiration}, nil + } + return &model.Link{ + RangeReader: stream.RangeReaderFunc(func(context.Context, http_range.Range) (io.ReadCloser, error) { + return io.NopCloser(strings.NewReader("file")), nil + }), + Expiration: &expiration, + }, nil +} + +func TestLinkCacheSeparatesRedirectAndProxy(t *testing.T) { + for _, tc := range []struct { + name string + firstRedirect bool + }{ + {name: "redirect then proxy", firstRedirect: true}, + {name: "proxy then redirect", firstRedirect: false}, + } { + t.Run(tc.name, func(t *testing.T) { + d := &linkModeDriver{storage: model.Storage{MountPath: "/" + t.Name()}} + for _, redirect := range []bool{tc.firstRedirect, !tc.firstRedirect, tc.firstRedirect, !tc.firstRedirect} { + link, _, err := Link(context.Background(), d, "/file", model.LinkArgs{Redirect: redirect}) + if err != nil { + t.Fatal(err) + } + if redirect && (link.URL == "" || link.RangeReader != nil) { + t.Fatalf("redirect link has wrong shape: %+v", link) + } + if !redirect && (link.URL != "" || link.RangeReader == nil) { + t.Fatalf("proxy link has wrong shape: %+v", link) + } + } + if d.calls != 2 { + t.Fatalf("expected one driver call per mode, got %d", d.calls) + } + }) + } +} diff --git a/server/common/proxy.go b/server/common/proxy.go index c76f43fc7..3522fe997 100644 --- a/server/common/proxy.go +++ b/server/common/proxy.go @@ -34,9 +34,7 @@ func Proxy(w http.ResponseWriter, r *http.Request, link *model.Link, file model. 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, &model.RangeReadCloser{ - RangeReader: rrf, - }) + return net.ServeHTTP(w, r, file.GetName(), file.ModTime(), size, rrf) } if link.RangeReader != nil { @@ -45,9 +43,7 @@ func Proxy(w http.ResponseWriter, r *http.Request, link *model.Link, file model. if size <= 0 { size = file.GetSize() } - return net.ServeHTTP(w, r, file.GetName(), file.ModTime(), size, &model.RangeReadCloser{ - RangeReader: link.RangeReader, - }) + return net.ServeHTTP(w, r, file.GetName(), file.ModTime(), size, link.RangeReader) } //transparent proxy diff --git a/server/s3/backend.go b/server/s3/backend.go index 779d98a60..de6413ce9 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 000000000..85df7fe81 --- /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 0191033d2..6624a07fc 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) {