mirror of
https://github.com/Wei-Shaw/sub2api.git
synced 2026-09-24 16:05:44 +08:00
Merge pull request #3716 from wucm667/fix/antigravity-refresh-server-invalidated-token
fix(antigravity): 401 服务端失效的 token 触发主动刷新,修复 NeedsRefresh 仅看 expires_at 的死循环
This commit is contained in:
@@ -12,6 +12,10 @@ const (
|
||||
// antigravityRefreshWindow Antigravity token 提前刷新窗口:15分钟
|
||||
// Google OAuth token 有效期55分钟,提前15分钟刷新
|
||||
antigravityRefreshWindow = 15 * time.Minute
|
||||
|
||||
antigravityForceTokenRefreshExtraKey = "antigravity_force_token_refresh"
|
||||
antigravityForceTokenRefreshReasonExtraKey = "antigravity_force_token_refresh_reason"
|
||||
antigravityForceTokenRefreshAtExtraKey = "antigravity_force_token_refresh_at"
|
||||
)
|
||||
|
||||
// AntigravityTokenRefresher 实现 TokenRefresher 接口
|
||||
@@ -41,6 +45,9 @@ func (r *AntigravityTokenRefresher) NeedsRefresh(account *Account, _ time.Durati
|
||||
if !r.CanRefresh(account) {
|
||||
return false
|
||||
}
|
||||
if accountNeedsAntigravityForceTokenRefresh(account) {
|
||||
return true
|
||||
}
|
||||
expiresAt := account.GetCredentialAsTime("expires_at")
|
||||
if expiresAt == nil {
|
||||
return false
|
||||
@@ -54,6 +61,29 @@ func (r *AntigravityTokenRefresher) NeedsRefresh(account *Account, _ time.Durati
|
||||
return needsRefresh
|
||||
}
|
||||
|
||||
func accountNeedsAntigravityForceTokenRefresh(account *Account) bool {
|
||||
return account != nil &&
|
||||
account.Platform == PlatformAntigravity &&
|
||||
account.Type == AccountTypeOAuth &&
|
||||
account.getExtraBool(antigravityForceTokenRefreshExtraKey)
|
||||
}
|
||||
|
||||
func antigravityForceTokenRefreshExtra(reason string) map[string]any {
|
||||
return map[string]any{
|
||||
antigravityForceTokenRefreshExtraKey: true,
|
||||
antigravityForceTokenRefreshReasonExtraKey: reason,
|
||||
antigravityForceTokenRefreshAtExtraKey: time.Now().UTC().Format(time.RFC3339),
|
||||
}
|
||||
}
|
||||
|
||||
func clearAntigravityForceTokenRefreshExtra() map[string]any {
|
||||
return map[string]any{
|
||||
antigravityForceTokenRefreshExtraKey: false,
|
||||
antigravityForceTokenRefreshReasonExtraKey: "",
|
||||
antigravityForceTokenRefreshAtExtraKey: "",
|
||||
}
|
||||
}
|
||||
|
||||
// Refresh 执行 token 刷新
|
||||
func (r *AntigravityTokenRefresher) Refresh(ctx context.Context, account *Account) (map[string]any, error) {
|
||||
tokenInfo, err := r.antigravityOAuthService.RefreshAccountToken(ctx, account)
|
||||
|
||||
@@ -292,6 +292,20 @@ func (s *RateLimitService) HandleUpstreamError(ctx context.Context, account *Acc
|
||||
if upstreamMsg != "" {
|
||||
msg = "OAuth 401: " + upstreamMsg
|
||||
}
|
||||
if authAccount.Platform == PlatformAntigravity {
|
||||
extraUpdates := antigravityForceTokenRefreshExtra("401_invalid")
|
||||
if err := s.accountRepo.UpdateExtra(ctx, authAccount.ID, extraUpdates); err != nil {
|
||||
slog.Warn("antigravity_401_force_refresh_mark_failed", "account_id", authAccount.ID, "error", err)
|
||||
} else {
|
||||
if authAccount.Extra == nil {
|
||||
authAccount.Extra = make(map[string]any, len(extraUpdates))
|
||||
}
|
||||
for k, v := range extraUpdates {
|
||||
authAccount.Extra[k] = v
|
||||
}
|
||||
slog.Info("antigravity_401_force_refresh_marked", "account_id", authAccount.ID)
|
||||
}
|
||||
}
|
||||
cooldownMinutes := s.cfg.RateLimit.OAuth401CooldownMinutes
|
||||
if cooldownMinutes <= 0 {
|
||||
cooldownMinutes = 10
|
||||
|
||||
@@ -18,7 +18,9 @@ type rateLimitAccountRepoStub struct {
|
||||
setErrorCalls int
|
||||
tempCalls int
|
||||
updateCredentialsCalls int
|
||||
updateExtraCalls int
|
||||
lastCredentials map[string]any
|
||||
lastExtraUpdates map[string]any
|
||||
lastErrorMsg string
|
||||
lastTempReason string
|
||||
lastErrorID int64
|
||||
@@ -45,6 +47,12 @@ func (r *rateLimitAccountRepoStub) UpdateCredentials(ctx context.Context, id int
|
||||
return nil
|
||||
}
|
||||
|
||||
func (r *rateLimitAccountRepoStub) UpdateExtra(ctx context.Context, id int64, updates map[string]any) error {
|
||||
r.updateExtraCalls++
|
||||
r.lastExtraUpdates = shallowCopyMap(updates)
|
||||
return nil
|
||||
}
|
||||
|
||||
type tokenCacheInvalidatorRecorder struct {
|
||||
accounts []*Account
|
||||
err error
|
||||
@@ -133,6 +141,10 @@ func TestRateLimitService_HandleUpstreamError_OAuth401SetsTempUnschedulable(t *t
|
||||
require.Equal(t, 1, repo.tempCalls)
|
||||
require.Equal(t, int64(100), repo.lastTempID)
|
||||
require.Contains(t, repo.lastTempReason, "invalid or expired credentials")
|
||||
require.Equal(t, 1, repo.updateExtraCalls)
|
||||
require.Equal(t, true, repo.lastExtraUpdates[antigravityForceTokenRefreshExtraKey])
|
||||
require.Equal(t, "401_invalid", repo.lastExtraUpdates[antigravityForceTokenRefreshReasonExtraKey])
|
||||
require.Equal(t, true, account.Extra[antigravityForceTokenRefreshExtraKey])
|
||||
require.Len(t, invalidator.accounts, 1)
|
||||
require.Equal(t, int64(100), invalidator.accounts[0].ID)
|
||||
})
|
||||
@@ -245,6 +257,7 @@ func TestRateLimitService_HandleUpstreamError_OAuth401DoesNotOverwriteCredential
|
||||
|
||||
require.True(t, shouldDisable)
|
||||
require.Equal(t, 0, repo.updateCredentialsCalls, "401 handler must not write credentials back from the request-start snapshot")
|
||||
require.Equal(t, 0, repo.updateExtraCalls, "OpenAI 401 must not set Antigravity force-refresh marker")
|
||||
require.Equal(t, 1, repo.tempCalls, "401 handler should still set temp-unschedulable cooldown")
|
||||
require.Nil(t, repo.lastCredentials, "no credentials should have been persisted")
|
||||
}
|
||||
|
||||
@@ -312,6 +312,7 @@ func (s *TokenRefreshService) refreshWithRetry(ctx context.Context, account *Acc
|
||||
if isNonRetryableRefreshError(err) {
|
||||
errorMsg := "Token refresh failed (non-retryable): " + logredact.RedactText(err.Error())
|
||||
s.notifyAccountSchedulingBlocked(account, time.Time{}, "token_refresh_non_retryable")
|
||||
s.clearAntigravityForceTokenRefresh(ctx, account, "non_retryable")
|
||||
if setErr := s.accountRepo.SetError(ctx, account.ID, errorMsg); setErr != nil {
|
||||
slog.Error("token_refresh.set_error_status_failed",
|
||||
"account_id", account.ID,
|
||||
@@ -369,6 +370,8 @@ func (s *TokenRefreshService) refreshWithRetry(ctx context.Context, account *Acc
|
||||
|
||||
// postRefreshActions 刷新成功后的后续动作(清除错误状态、缓存失效、调度器同步等)
|
||||
func (s *TokenRefreshService) postRefreshActions(ctx context.Context, account *Account) {
|
||||
s.clearAntigravityForceTokenRefresh(ctx, account, "success")
|
||||
|
||||
// Antigravity 账户:如果之前是因为缺少 project_id 而标记为 error,现在成功获取到了,清除错误状态
|
||||
if account.Platform == PlatformAntigravity &&
|
||||
account.Status == StatusError &&
|
||||
@@ -432,6 +435,30 @@ func (s *TokenRefreshService) postRefreshActions(ctx context.Context, account *A
|
||||
s.ensureAntigravityPrivacy(ctx, account)
|
||||
}
|
||||
|
||||
func (s *TokenRefreshService) clearAntigravityForceTokenRefresh(ctx context.Context, account *Account, outcome string) {
|
||||
if s == nil || account == nil || !accountNeedsAntigravityForceTokenRefresh(account) {
|
||||
return
|
||||
}
|
||||
updates := clearAntigravityForceTokenRefreshExtra()
|
||||
if err := s.accountRepo.UpdateExtra(ctx, account.ID, updates); err != nil {
|
||||
slog.Warn("token_refresh.clear_antigravity_force_refresh_failed",
|
||||
"account_id", account.ID,
|
||||
"outcome", outcome,
|
||||
"error", err,
|
||||
)
|
||||
return
|
||||
}
|
||||
if account.Extra != nil {
|
||||
for k, v := range updates {
|
||||
account.Extra[k] = v
|
||||
}
|
||||
}
|
||||
slog.Info("token_refresh.cleared_antigravity_force_refresh",
|
||||
"account_id", account.ID,
|
||||
"outcome", outcome,
|
||||
)
|
||||
}
|
||||
|
||||
// errRefreshSkipped 表示刷新被跳过(锁竞争或已被其他路径刷新),不计入 failed 或 refreshed
|
||||
var errRefreshSkipped = fmt.Errorf("refresh skipped")
|
||||
|
||||
|
||||
@@ -20,8 +20,10 @@ type tokenRefreshAccountRepo struct {
|
||||
setErrorCalls int
|
||||
clearTempCalls int
|
||||
setTempUnschedCalls int
|
||||
updateExtraCalls int
|
||||
lastErrorMessage string
|
||||
lastTempUnschedReason string
|
||||
lastExtraUpdates map[string]any
|
||||
lastAccount *Account
|
||||
updateErr error
|
||||
}
|
||||
@@ -68,6 +70,22 @@ func (r *tokenRefreshAccountRepo) SetTempUnschedulable(ctx context.Context, id i
|
||||
return nil
|
||||
}
|
||||
|
||||
func (r *tokenRefreshAccountRepo) UpdateExtra(ctx context.Context, id int64, updates map[string]any) error {
|
||||
r.updateExtraCalls++
|
||||
r.lastExtraUpdates = shallowCopyMap(updates)
|
||||
if r.accountsByID != nil {
|
||||
if acc, ok := r.accountsByID[id]; ok && acc != nil {
|
||||
if acc.Extra == nil {
|
||||
acc.Extra = make(map[string]any, len(updates))
|
||||
}
|
||||
for k, v := range updates {
|
||||
acc.Extra[k] = v
|
||||
}
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
type tokenCacheInvalidatorStub struct {
|
||||
calls int
|
||||
err error
|
||||
@@ -233,6 +251,121 @@ func TestTokenRefreshService_RefreshWithRetry_Antigravity(t *testing.T) {
|
||||
require.Equal(t, 1, invalidator.calls) // Antigravity 也应触发缓存失效
|
||||
}
|
||||
|
||||
func TestAntigravityTokenRefresher_NeedsRefresh_ForceRefreshMarker(t *testing.T) {
|
||||
refresher := NewAntigravityTokenRefresher(nil)
|
||||
account := &Account{
|
||||
ID: 3675,
|
||||
Platform: PlatformAntigravity,
|
||||
Type: AccountTypeOAuth,
|
||||
Credentials: map[string]any{
|
||||
"expires_at": time.Now().Add(time.Hour).Format(time.RFC3339),
|
||||
},
|
||||
Extra: map[string]any{
|
||||
antigravityForceTokenRefreshExtraKey: true,
|
||||
},
|
||||
}
|
||||
|
||||
require.True(t, refresher.NeedsRefresh(account, 0), "server-invalidated token must refresh even before expires_at")
|
||||
}
|
||||
|
||||
func TestAntigravityTokenRefresher_NeedsRefresh_NormalExpiryRulesUnchanged(t *testing.T) {
|
||||
refresher := NewAntigravityTokenRefresher(nil)
|
||||
|
||||
t.Run("normal_unexpired_without_marker_does_not_refresh", func(t *testing.T) {
|
||||
account := &Account{
|
||||
ID: 3707,
|
||||
Platform: PlatformAntigravity,
|
||||
Type: AccountTypeOAuth,
|
||||
Credentials: map[string]any{
|
||||
"expires_at": time.Now().Add(time.Hour).Format(time.RFC3339),
|
||||
},
|
||||
}
|
||||
|
||||
require.False(t, refresher.NeedsRefresh(account, 0))
|
||||
})
|
||||
|
||||
t.Run("normal_expiring_refreshes", func(t *testing.T) {
|
||||
account := &Account{
|
||||
ID: 3708,
|
||||
Platform: PlatformAntigravity,
|
||||
Type: AccountTypeOAuth,
|
||||
Credentials: map[string]any{
|
||||
"expires_at": time.Now().Add(5 * time.Minute).Format(time.RFC3339),
|
||||
},
|
||||
}
|
||||
|
||||
require.True(t, refresher.NeedsRefresh(account, 0))
|
||||
})
|
||||
}
|
||||
|
||||
func TestTokenRefreshService_RefreshWithRetry_AntigravityClearsForceRefreshOnSuccess(t *testing.T) {
|
||||
repo := &tokenRefreshAccountRepo{}
|
||||
cfg := &config.Config{
|
||||
TokenRefresh: config.TokenRefreshConfig{
|
||||
MaxRetries: 1,
|
||||
RetryBackoffSeconds: 0,
|
||||
},
|
||||
}
|
||||
service := NewTokenRefreshService(repo, nil, nil, nil, nil, nil, nil, cfg, nil)
|
||||
until := time.Now().Add(10 * time.Minute)
|
||||
account := &Account{
|
||||
ID: 3709,
|
||||
Platform: PlatformAntigravity,
|
||||
Type: AccountTypeOAuth,
|
||||
TempUnschedulableUntil: &until,
|
||||
Extra: map[string]any{
|
||||
antigravityForceTokenRefreshExtraKey: true,
|
||||
antigravityForceTokenRefreshReasonExtraKey: "401_invalid",
|
||||
"privacy_mode": AntigravityPrivacySet,
|
||||
},
|
||||
}
|
||||
refresher := &tokenRefresherStub{
|
||||
credentials: map[string]any{
|
||||
"access_token": "new-ag-token",
|
||||
},
|
||||
}
|
||||
|
||||
err := service.refreshWithRetry(context.Background(), account, refresher, refresher, time.Hour)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, 1, repo.updateCredentialsCalls)
|
||||
require.Equal(t, 1, repo.updateExtraCalls)
|
||||
require.Equal(t, false, repo.lastExtraUpdates[antigravityForceTokenRefreshExtraKey])
|
||||
require.Equal(t, "", repo.lastExtraUpdates[antigravityForceTokenRefreshReasonExtraKey])
|
||||
require.Equal(t, false, account.Extra[antigravityForceTokenRefreshExtraKey])
|
||||
require.Equal(t, 1, repo.clearTempCalls, "successful refresh should restore schedulability")
|
||||
}
|
||||
|
||||
func TestTokenRefreshService_RefreshWithRetry_AntigravityForceRefreshInvalidGrantSetsError(t *testing.T) {
|
||||
repo := &tokenRefreshAccountRepo{}
|
||||
cfg := &config.Config{
|
||||
TokenRefresh: config.TokenRefreshConfig{
|
||||
MaxRetries: 3,
|
||||
RetryBackoffSeconds: 0,
|
||||
},
|
||||
}
|
||||
service := NewTokenRefreshService(repo, nil, nil, nil, nil, nil, nil, cfg, nil)
|
||||
account := &Account{
|
||||
ID: 3710,
|
||||
Platform: PlatformAntigravity,
|
||||
Type: AccountTypeOAuth,
|
||||
Extra: map[string]any{
|
||||
antigravityForceTokenRefreshExtraKey: true,
|
||||
antigravityForceTokenRefreshReasonExtraKey: "401_invalid",
|
||||
},
|
||||
}
|
||||
refresher := &tokenRefresherStub{
|
||||
err: errors.New("invalid_grant: token revoked"),
|
||||
}
|
||||
|
||||
err := service.refreshWithRetry(context.Background(), account, refresher, refresher, time.Hour)
|
||||
require.Error(t, err)
|
||||
require.Equal(t, 1, repo.setErrorCalls)
|
||||
require.Equal(t, 0, repo.setTempUnschedCalls)
|
||||
require.Equal(t, 1, repo.updateExtraCalls)
|
||||
require.Equal(t, false, repo.lastExtraUpdates[antigravityForceTokenRefreshExtraKey])
|
||||
require.Contains(t, repo.lastErrorMessage, "non-retryable")
|
||||
}
|
||||
|
||||
// TestTokenRefreshService_RefreshWithRetry_NonOAuthAccount 测试非 OAuth 账号不触发缓存失效
|
||||
func TestTokenRefreshService_RefreshWithRetry_NonOAuthAccount(t *testing.T) {
|
||||
repo := &tokenRefreshAccountRepo{}
|
||||
|
||||
Reference in New Issue
Block a user