mirror of
https://github.com/krau/SaveAny-Bot.git
synced 2026-08-19 11:23:57 +08:00
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:
@@ -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)
|
||||||
|
|||||||
@@ -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
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user