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 <nostalucent@gmail.com>
Co-authored-by: Codex <267193182+codex@users.noreply.github.com>
This commit is contained in:
Nostalgia
2026-09-24 12:01:50 +08:00
committed by GitHub
parent 90acfa18e4
commit 893457cd50
16 changed files with 1139 additions and 69 deletions
+295
View File
@@ -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, "<redacted>")
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
}
+365
View File
@@ -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())
}
}
+16 -2
View File
@@ -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"
@@ -24,6 +25,7 @@ type AliyundriveOpen struct {
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) {
+1
View File
@@ -20,6 +20,7 @@ type Addition struct {
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
}
+1 -1
View File
@@ -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
+2 -2
View File
@@ -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=
+1
View File
@@ -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")
-18
View File
@@ -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
}
+54 -24
View File
@@ -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
+236
View File
@@ -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() }
+4 -1
View File
@@ -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
}
+71
View File
@@ -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)
}
})
}
}
+2 -6
View File
@@ -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
+2
View File
@@ -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
+66
View File
@@ -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)
}
}
+8
View File
@@ -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) {