Merge pull request #3716 from wucm667/fix/antigravity-refresh-server-invalidated-token

fix(antigravity): 401 服务端失效的 token 触发主动刷新,修复 NeedsRefresh 仅看 expires_at 的死循环
This commit is contained in:
Wesley Liddick
2026-07-06 10:18:49 +08:00
committed by GitHub
5 changed files with 217 additions and 0 deletions
@@ -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{}