mirror of
https://github.com/OpenListTeam/OpenList.git
synced 2026-10-10 04:53:09 +08:00
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:
@@ -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
|
||||
}
|
||||
@@ -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())
|
||||
}
|
||||
}
|
||||
@@ -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) {
|
||||
|
||||
@@ -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{
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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=
|
||||
|
||||
@@ -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")
|
||||
|
||||
@@ -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
@@ -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
|
||||
|
||||
@@ -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
@@ -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
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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) {
|
||||
|
||||
Reference in New Issue
Block a user