Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
46 changes: 46 additions & 0 deletions backend/internal/application/account/credential_refresh_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -406,6 +406,52 @@ func TestCredentialRefreshFailureDistinguishesTransientAndPermanent(t *testing.T
}
}

func TestCredentialDecryptFailedAllowsRetryAfterKeyRecovery(t *testing.T) {
ctx := context.Background()
now := time.Now().UTC()
service, credential, adapter := newCredentialRefreshTestService(t, now)
service.now = func() time.Time { return now }

// 旧行为会把 decrypt_failed 标 permanent;模拟已落库的 permanent 状态。
if err := service.accounts.UpdateCredentialRefreshFailure(ctx, credential.ID, 1, now.Add(time.Hour), "credential_decrypt_failed", true); err != nil {
t.Fatal(err)
}
stuck, err := service.accounts.Get(ctx, credential.ID)
if err != nil || !stuck.RefreshPermanent || stuck.LastRefreshErrorCode != "credential_decrypt_failed" {
t.Fatalf("setup stuck state = %#v err=%v", stuck, err)
}

// 密钥恢复后:手动 force 必须能再次发起刷新。
adapter.refreshErr = nil
service.clearRefreshState(credential.ID)
recovered, err := service.EnsureCredential(ctx, stuck, true)
if err != nil {
t.Fatalf("force refresh after decrypt_failed should retry: %v", err)
}
if recovered.RefreshPermanent || recovered.LastRefreshErrorCode != "" || adapter.refreshCount.Load() < 1 {
t.Fatalf("decrypt_failed was not cleared after successful refresh: %#v count=%d", recovered, adapter.refreshCount.Load())
}

// invalid_grant 仍须保持永久阻断。
service.clearRefreshState(credential.ID)
adapter.refreshErr = &provider.CredentialRefreshError{Status: 400, Code: "invalid_grant", Permanent: true}
if _, err := service.EnsureCredential(ctx, recovered, true); err == nil {
t.Fatal("invalid_grant should fail")
}
blocked, err := service.accounts.Get(ctx, credential.ID)
if err != nil || !blocked.RefreshPermanent || blocked.LastRefreshErrorCode != "invalid_grant" {
t.Fatalf("invalid_grant permanent state = %#v err=%v", blocked, err)
}
// force 也不得再打 OAuth(真正永久)
count := adapter.refreshCount.Load()
if _, err := service.EnsureCredential(ctx, blocked, true); err == nil {
t.Fatal("invalid_grant force should still be blocked")
}
if adapter.refreshCount.Load() != count {
t.Fatalf("invalid_grant forced another oauth call: before=%d after=%d", count, adapter.refreshCount.Load())
}
}

func TestRefreshAllTokensSkipsUnrefreshableAccounts(t *testing.T) {
ctx := context.Background()
now := time.Date(2026, 7, 11, 12, 0, 0, 0, time.UTC)
Expand Down
4 changes: 2 additions & 2 deletions backend/internal/application/account/credential_scheduler.go
Original file line number Diff line number Diff line change
Expand Up @@ -56,7 +56,7 @@ func (s *Service) RecoverCriticalCredentials(ctx context.Context, expiresWithin
if getErr != nil {
return getErr
}
if credential.RefreshPermanent {
if credential.RefreshPermanent && !isRecoverableRefreshErrorCode(credential.LastRefreshErrorCode) {
if !credential.ExpiresAt.IsZero() && credential.ExpiresAt.After(s.now()) {
return nil
}
Expand Down Expand Up @@ -130,7 +130,7 @@ func (s *Service) refreshDueCredentials(ctx context.Context) error {
if !credential.Enabled || credential.AuthStatus != accountdomain.AuthStatusActive || s.providers == nil || !s.providers.SupportsCredentialRefresh(credential.Provider) || credential.EncryptedRefreshToken == "" {
return nil
}
if credential.RefreshPermanent {
if credential.RefreshPermanent && !isRecoverableRefreshErrorCode(credential.LastRefreshErrorCode) {
if !credential.ExpiresAt.IsZero() && credential.ExpiresAt.After(s.now()) {
return nil
}
Expand Down
26 changes: 24 additions & 2 deletions backend/internal/application/account/service.go
Original file line number Diff line number Diff line change
Expand Up @@ -1744,8 +1744,14 @@ func (s *Service) recordCredentialRefreshFailure(ctx context.Context, credential
} else if errors.Is(refreshErr, context.DeadlineExceeded) {
errorCode = "oauth_timeout"
}
// 永久失败只能由成功换取新 token 清除,后续偶发传输错误不能把状态降级为可重试。
permanent = permanent || credential.RefreshPermanent
// 真正的 OAuth 永久失败(invalid_grant 等)只能由成功换 token 清除。
// credential_decrypt_failed 是可恢复本地错误:不得被旧 permanent 粘住,也不得把本次可恢复失败抬升为永久。
if permanent && isRecoverableRefreshErrorCode(errorCode) {
permanent = false
}
if credential.RefreshPermanent && !isRecoverableRefreshErrorCode(credential.LastRefreshErrorCode) && !isRecoverableRefreshErrorCode(errorCode) {
permanent = true
}
now := s.now()
retryAt := now.Add(credentialRefreshBackoff(credential.ID, failureCount, retryAfter))
accessTokenAlive := credential.EncryptedAccessToken != "" && !credential.ExpiresAt.IsZero() && credential.ExpiresAt.After(now)
Expand Down Expand Up @@ -1774,10 +1780,16 @@ func (s *Service) recordCredentialRefreshFailure(ctx context.Context, credential
}

// resolvePermanentRefreshFailure 阻止再次请求已确认失效的 refresh token,并在 access token 到期后收敛账号状态。
// credential_decrypt_failed 属于本地密钥问题,允许手动 force / 调度重试(密钥恢复后可自愈);
// invalid_grant 等真正 OAuth 永久失败仍保持阻断。
func (s *Service) resolvePermanentRefreshFailure(ctx context.Context, credential accountdomain.Credential, now time.Time, force bool) (accountdomain.Credential, error, bool) {
if !credential.RefreshPermanent {
return accountdomain.Credential{}, nil, false
}
if isRecoverableRefreshErrorCode(credential.LastRefreshErrorCode) {
// 允许 force 或到期调度再次尝试解密/刷新;成功后会 clear permanent 标记。
return accountdomain.Credential{}, nil, false
}
accessTokenAlive := credential.EncryptedAccessToken != "" && !credential.ExpiresAt.IsZero() && credential.ExpiresAt.After(now)
if accessTokenAlive && !force {
return credential, nil, true
Expand All @@ -1793,6 +1805,16 @@ func (s *Service) resolvePermanentRefreshFailure(ctx context.Context, credential
return accountdomain.Credential{}, fmt.Errorf("%w: %s", ErrCredentialRefreshPermanent, credential.LastRefreshErrorCode), true
}

// isRecoverableRefreshErrorCode 标识“永久标记可被后续成功刷新清除”的本地/临时错误。
func isRecoverableRefreshErrorCode(code string) bool {
switch strings.TrimSpace(code) {
case "credential_decrypt_failed":
return true
default:
return false
}
}

func credentialRefreshBackoff(accountID uint64, failureCount int, retryAfter time.Duration) time.Duration {
delays := [...]time.Duration{30 * time.Second, 2 * time.Minute, 5 * time.Minute, 10 * time.Minute, 15 * time.Minute}
index := max(0, min(failureCount-1, len(delays)-1))
Expand Down
61 changes: 61 additions & 0 deletions backend/internal/application/gateway/official_cache_e2e_test.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,61 @@
package gateway

import (
"testing"

accountdomain "github.com/chenyme/grok2api/backend/internal/domain/account"
)

// 模拟 Sub2API → Grok2API 官方缓存亲和链路(不连真上游)。
func TestOfficialCache_Sub2APIPromptCacheKeyStickyAcrossTurns(t *testing.T) {
// Sub2API 透传 body.prompt_cache_key
explicit := "sub2api-stable-session-uuid-001"
turn1Body := []byte(`{"model":"grok-4.5","prompt_cache_key":"sub2api-stable-session-uuid-001","messages":[{"role":"user","content":"hello"}]}`)
turn2Body := []byte(`{"model":"grok-4.5","prompt_cache_key":"sub2api-stable-session-uuid-001","messages":[{"role":"user","content":"hello"},{"role":"assistant","content":"hi"},{"role":"user","content":"again"}]}`)

id1 := resolveBuildSessionIdentity(42, accountdomain.ProviderBuild, "grok-4.5", explicit, "", turn1Body)
id2 := resolveBuildSessionIdentity(42, accountdomain.ProviderBuild, "grok-4.5", explicit, "", turn2Body)
if id1.upstreamID == "" || id1.upstreamID != id2.upstreamID {
t.Fatalf("sub2api prompt_cache_key must yield stable upstream id: %#v %#v", id1, id2)
}
if id1.affinityKey == "" || id1.affinityKey != id2.affinityKey {
t.Fatalf("affinity must stick across turns: %#v %#v", id1, id2)
}
// 不同租户(clientKey)必须隔离
other := resolveBuildSessionIdentity(99, accountdomain.ProviderBuild, "grok-4.5", explicit, "", turn1Body)
if other.upstreamID == id1.upstreamID {
t.Fatal("tenant isolation broken for explicit prompt_cache_key")
}
}

func TestOfficialCache_Sub2APIMissingSessionFallsBackToMessageHash(t *testing.T) {
// Sub2API 没透传任何 session:靠首条 user 内容 soft 粘滞
turn1 := []byte(`{"messages":[{"role":"system","content":"sys"},{"role":"user","content":"what is mutex"}]}`)
turn2 := []byte(`{"messages":[{"role":"system","content":"sys"},{"role":"user","content":"what is mutex"},{"role":"assistant","content":"A lock."},{"role":"user","content":"and deadlock?"}]}`)
id1 := resolveBuildSessionIdentity(7, accountdomain.ProviderBuild, "grok-4.5", "", "", turn1)
id2 := resolveBuildSessionIdentity(7, accountdomain.ProviderBuild, "grok-4.5", "", "", turn2)
if !id1.soft || id1.upstreamID == "" {
t.Fatalf("expected soft session, got %#v", id1)
}
if id1.upstreamID != id2.upstreamID {
t.Fatalf("soft session drifted: %s vs %s", id1.upstreamID, id2.upstreamID)
}
}

func TestOfficialCache_NoSignalMeansNoRandomConvID(t *testing.T) {
// 完全无信号:不得生成上游 ID(由 adapter 侧保证不写随机 conv-id)
id := resolveBuildSessionIdentity(7, accountdomain.ProviderBuild, "grok-4.5", "", "", []byte(`{}`))
if id.upstreamID != "" || id.affinityKey != "" {
t.Fatalf("empty signal must not invent session: %#v", id)
}
}

func TestOfficialCache_ClaudeCodeSessionHeaderPath(t *testing.T) {
// 模拟 handler 已把 X-Claude-Code-Session-Id 抽成 seed
seed := "123e4567-e89b-12d3-a456-426614174000"
a := resolveBuildSessionIdentity(1, accountdomain.ProviderBuild, "grok-4.5", "", seed, nil)
b := resolveBuildSessionIdentity(1, accountdomain.ProviderBuild, "grok-4.5", "", seed, []byte(`{"messages":[{"role":"user","content":"x"}]}`))
if a.upstreamID == "" || a.upstreamID != b.upstreamID {
t.Fatalf("claude session seed unstable: %#v %#v", a, b)
}
}
Loading
Loading