From b23475ac0c553089d37adaf53f750c54fa0224e4 Mon Sep 17 00:00:00 2001 From: wucm667 Date: Sun, 5 Jul 2026 08:34:34 +0800 Subject: [PATCH] fix(antigravity): refresh server-invalidated tokens --- .../service/antigravity_token_refresher.go | 30 ++++ backend/internal/service/ratelimit_service.go | 14 ++ .../service/ratelimit_service_401_test.go | 13 ++ .../internal/service/token_refresh_service.go | 27 ++++ .../service/token_refresh_service_test.go | 133 ++++++++++++++++++ 5 files changed, 217 insertions(+) diff --git a/backend/internal/service/antigravity_token_refresher.go b/backend/internal/service/antigravity_token_refresher.go index 7ce0ccf0fe..f40df7d95f 100644 --- a/backend/internal/service/antigravity_token_refresher.go +++ b/backend/internal/service/antigravity_token_refresher.go @@ -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) diff --git a/backend/internal/service/ratelimit_service.go b/backend/internal/service/ratelimit_service.go index 2b9c85a2fa..81a915be67 100644 --- a/backend/internal/service/ratelimit_service.go +++ b/backend/internal/service/ratelimit_service.go @@ -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 diff --git a/backend/internal/service/ratelimit_service_401_test.go b/backend/internal/service/ratelimit_service_401_test.go index 09afb9314b..48e6a41def 100644 --- a/backend/internal/service/ratelimit_service_401_test.go +++ b/backend/internal/service/ratelimit_service_401_test.go @@ -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") } diff --git a/backend/internal/service/token_refresh_service.go b/backend/internal/service/token_refresh_service.go index 9ca369a52a..2d4a026692 100644 --- a/backend/internal/service/token_refresh_service.go +++ b/backend/internal/service/token_refresh_service.go @@ -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") diff --git a/backend/internal/service/token_refresh_service_test.go b/backend/internal/service/token_refresh_service_test.go index 6f227352f9..bd6a521576 100644 --- a/backend/internal/service/token_refresh_service_test.go +++ b/backend/internal/service/token_refresh_service_test.go @@ -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{}