fix: deduplicate alist token refreshes

Guard token access with a mutex and merge concurrent logins.
Reuse a recent refresh to avoid login storms.
This commit is contained in:
krau
2026-08-16 21:20:11 +08:00
parent 534ed7a7c2
commit 189bf9c736
2 changed files with 54 additions and 11 deletions

View File

@@ -9,6 +9,7 @@ import (
"net/http" "net/http"
"net/url" "net/url"
"path" "path"
"sync"
"time" "time"
"github.com/charmbracelet/log" "github.com/charmbracelet/log"
@@ -17,15 +18,26 @@ import (
"github.com/krau/SaveAny-Bot/pkg/enums/ctxkey" "github.com/krau/SaveAny-Bot/pkg/enums/ctxkey"
storenum "github.com/krau/SaveAny-Bot/pkg/enums/storage" storenum "github.com/krau/SaveAny-Bot/pkg/enums/storage"
"github.com/krau/SaveAny-Bot/pkg/storagetypes" "github.com/krau/SaveAny-Bot/pkg/storagetypes"
"golang.org/x/sync/singleflight"
) )
type Alist struct { type Alist struct {
client *http.Client client *http.Client
token string tokenMu sync.RWMutex
baseURL string token string
loginInfo *loginRequest lastLoginAt time.Time
config config.AlistStorageConfig tokenFlight singleflight.Group
logger *log.Logger baseURL string
loginInfo *loginRequest
config config.AlistStorageConfig
logger *log.Logger
}
// authHeader returns the current token for use in API requests.
func (a *Alist) authHeader() string {
a.tokenMu.RLock()
defer a.tokenMu.RUnlock()
return a.token
} }
func (a *Alist) Init(ctx context.Context, cfg config.StorageConfig) error { func (a *Alist) Init(ctx context.Context, cfg config.StorageConfig) error {
@@ -42,14 +54,16 @@ func (a *Alist) Init(ctx context.Context, cfg config.StorageConfig) error {
a.logger = log.FromContext(ctx).WithPrefix(fmt.Sprintf("alist[%s]", alistConfig.Name)) a.logger = log.FromContext(ctx).WithPrefix(fmt.Sprintf("alist[%s]", alistConfig.Name))
if alistConfig.Token != "" { if alistConfig.Token != "" {
a.tokenMu.Lock()
a.token = alistConfig.Token a.token = alistConfig.Token
a.tokenMu.Unlock()
tokenCtx, cancel := context.WithTimeout(ctx, 1*time.Minute) tokenCtx, cancel := context.WithTimeout(ctx, 1*time.Minute)
defer cancel() defer cancel()
req, err := http.NewRequestWithContext(tokenCtx, http.MethodGet, a.baseURL+"/api/me", nil) req, err := http.NewRequestWithContext(tokenCtx, http.MethodGet, a.baseURL+"/api/me", nil)
if err != nil { if err != nil {
return fmt.Errorf("failed to create request: %w", err) return fmt.Errorf("failed to create request: %w", err)
} }
req.Header.Set("Authorization", a.token) req.Header.Set("Authorization", a.authHeader())
resp, err := a.client.Do(req) resp, err := a.client.Do(req)
if err != nil { if err != nil {
@@ -157,7 +171,7 @@ func (a *Alist) putFile(ctx context.Context, reader io.Reader, storagePath strin
if err != nil { if err != nil {
return nil, fmt.Errorf("failed to create request: %w", err) return nil, fmt.Errorf("failed to create request: %w", err)
} }
req.Header.Set("Authorization", a.token) req.Header.Set("Authorization", a.authHeader())
req.Header.Set("File-Path", url.PathEscape(storagePath)) req.Header.Set("File-Path", url.PathEscape(storagePath))
req.Header.Set("Content-Type", "application/octet-stream") req.Header.Set("Content-Type", "application/octet-stream")
if length := ctx.Value(ctxkey.ContentLength); length != nil { if length := ctx.Value(ctxkey.ContentLength); length != nil {
@@ -208,7 +222,7 @@ func (a *Alist) existsPath(ctx context.Context, storagePath string) bool {
a.logger.Errorf("Failed to create request: %v", err) a.logger.Errorf("Failed to create request: %v", err)
return false return false
} }
req.Header.Set("Authorization", a.token) req.Header.Set("Authorization", a.authHeader())
req.Header.Set("Content-Type", "application/json") req.Header.Set("Content-Type", "application/json")
resp, err := a.client.Do(req) resp, err := a.client.Do(req)
if err != nil { if err != nil {
@@ -263,7 +277,7 @@ func (a *Alist) ListFiles(ctx context.Context, dirPath string) ([]storagetypes.F
if err != nil { if err != nil {
return nil, fmt.Errorf("failed to create request: %w", err) return nil, fmt.Errorf("failed to create request: %w", err)
} }
req.Header.Set("Authorization", a.token) req.Header.Set("Authorization", a.authHeader())
req.Header.Set("Content-Type", "application/json") req.Header.Set("Content-Type", "application/json")
resp, err := a.client.Do(req) resp, err := a.client.Do(req)
@@ -338,7 +352,7 @@ func (a *Alist) OpenFile(ctx context.Context, filePath string) (io.ReadCloser, i
if err != nil { if err != nil {
return nil, 0, fmt.Errorf("failed to create request: %w", err) return nil, 0, fmt.Errorf("failed to create request: %w", err)
} }
req.Header.Set("Authorization", a.token) req.Header.Set("Authorization", a.authHeader())
req.Header.Set("Content-Type", "application/json") req.Header.Set("Content-Type", "application/json")
resp, err := a.client.Do(req) resp, err := a.client.Do(req)

View File

@@ -12,7 +12,33 @@ import (
config "github.com/krau/SaveAny-Bot/config/storage" config "github.com/krau/SaveAny-Bot/config/storage"
) )
// minTokenRefreshInterval deduplicates login storms: a successful refresh is
// reused for this window instead of hitting the login endpoint again.
const minTokenRefreshInterval = 30 * time.Second
// getToken refreshes the JWT, deduplicating concurrent calls so parallel
// uploads and the background refresher share one login request.
func (a *Alist) getToken(ctx context.Context) error { func (a *Alist) getToken(ctx context.Context) error {
a.tokenMu.RLock()
fresh := !a.lastLoginAt.IsZero() && time.Since(a.lastLoginAt) < minTokenRefreshInterval
a.tokenMu.RUnlock()
if fresh {
return nil
}
_, err, _ := a.tokenFlight.Do("token", func() (any, error) {
// Another waiter may have refreshed while this call was queued.
a.tokenMu.RLock()
fresh := !a.lastLoginAt.IsZero() && time.Since(a.lastLoginAt) < minTokenRefreshInterval
a.tokenMu.RUnlock()
if fresh {
return nil, nil
}
return nil, a.fetchToken(ctx)
})
return err
}
func (a *Alist) fetchToken(ctx context.Context) error {
loginBody, err := json.Marshal(a.loginInfo) loginBody, err := json.Marshal(a.loginInfo)
if err != nil { if err != nil {
return fmt.Errorf("failed to marshal login request: %w", err) return fmt.Errorf("failed to marshal login request: %w", err)
@@ -44,7 +70,10 @@ func (a *Alist) getToken(ctx context.Context) error {
return fmt.Errorf("%w: %s", ErrAlistLoginFailed, loginResp.Message) return fmt.Errorf("%w: %s", ErrAlistLoginFailed, loginResp.Message)
} }
a.tokenMu.Lock()
a.token = loginResp.Data.Token a.token = loginResp.Data.Token
a.lastLoginAt = time.Now()
a.tokenMu.Unlock()
return nil return nil
} }