fix: reduce grok oauth account-risk paths

This commit is contained in:
Heatherm Huang
2026-06-26 10:36:09 +08:00
parent b2e2c7e69c
commit f29ccc7dfb
11 changed files with 124 additions and 25 deletions
+10 -7
View File
@@ -14,7 +14,7 @@ import (
)
const (
OAuthIssuer = "https://accounts.x.ai"
OAuthIssuer = "https://auth.x.ai"
DiscoveryURL = OAuthIssuer + "/.well-known/openid-configuration"
DefaultAuthorizeURL = OAuthIssuer + "/oauth2/authorize"
DefaultTokenURL = OAuthIssuer + "/oauth2/token"
@@ -217,8 +217,9 @@ func BuildAuthorizationURL(state, codeChallenge, redirectURI, nonce string) stri
// AuthorizationInput is a parsed manual OAuth callback input.
type AuthorizationInput struct {
Code string
State string
Code string
State string
RequiresState bool
}
// ParseAuthorizationInput accepts a full callback URL, query string, or bare code.
@@ -232,8 +233,9 @@ func ParseAuthorizationInput(raw string) AuthorizationInput {
values := parsed.Query()
if code := strings.TrimSpace(values.Get("code")); code != "" {
return AuthorizationInput{
Code: code,
State: strings.TrimSpace(values.Get("state")),
Code: code,
State: strings.TrimSpace(values.Get("state")),
RequiresState: true,
}
}
}
@@ -243,8 +245,9 @@ func ParseAuthorizationInput(raw string) AuthorizationInput {
if values, err := url.ParseQuery(queryCandidate); err == nil {
if code := strings.TrimSpace(values.Get("code")); code != "" {
return AuthorizationInput{
Code: code,
State: strings.TrimSpace(values.Get("state")),
Code: code,
State: strings.TrimSpace(values.Get("state")),
RequiresState: true,
}
}
}
+28 -12
View File
@@ -13,22 +13,37 @@ func TestParseAuthorizationInput(t *testing.T) {
t.Parallel()
tests := []struct {
name string
raw string
wantCode string
wantState string
name string
raw string
wantCode string
wantState string
wantRequiresState bool
}{
{
name: "full callback url",
raw: "http://127.0.0.1:56121/callback?code=abc123&state=state456",
wantCode: "abc123",
wantState: "state456",
name: "full callback url",
raw: "http://127.0.0.1:56121/callback?code=abc123&state=state456",
wantCode: "abc123",
wantState: "state456",
wantRequiresState: true,
},
{
name: "query string",
raw: "?code=abc123&state=state456",
wantCode: "abc123",
wantState: "state456",
name: "query string",
raw: "?code=abc123&state=state456",
wantCode: "abc123",
wantState: "state456",
wantRequiresState: true,
},
{
name: "full callback url missing state",
raw: "http://127.0.0.1:56121/callback?code=abc123",
wantCode: "abc123",
wantRequiresState: true,
},
{
name: "query string missing state",
raw: "code=abc123",
wantCode: "abc123",
wantRequiresState: true,
},
{
name: "bare code",
@@ -43,6 +58,7 @@ func TestParseAuthorizationInput(t *testing.T) {
got := ParseAuthorizationInput(tt.raw)
require.Equal(t, tt.wantCode, got.Code)
require.Equal(t, tt.wantState, got.State)
require.Equal(t, tt.wantRequiresState, got.RequiresState)
})
}
}
+9 -2
View File
@@ -2572,6 +2572,13 @@ func (s *adminServiceImpl) GetAccountsByIDs(ctx context.Context, ids []int64) ([
return accounts, nil
}
func normalizeAccountConcurrency(platform, accountType string, concurrency int) int {
if platform == PlatformGrok && accountType == AccountTypeOAuth && concurrency <= 0 {
return 1
}
return concurrency
}
func (s *adminServiceImpl) CreateAccount(ctx context.Context, input *CreateAccountInput) (*Account, error) {
// 绑定分组
groupIDs := input.GroupIDs
@@ -2604,7 +2611,7 @@ func (s *adminServiceImpl) CreateAccount(ctx context.Context, input *CreateAccou
Credentials: input.Credentials,
Extra: input.Extra,
ProxyID: input.ProxyID,
Concurrency: input.Concurrency,
Concurrency: normalizeAccountConcurrency(input.Platform, input.Type, input.Concurrency),
Priority: input.Priority,
Status: StatusActive,
Schedulable: true,
@@ -2737,7 +2744,7 @@ func (s *adminServiceImpl) UpdateAccount(ctx context.Context, id int64, input *U
}
// 只在指针非 nil 时更新 Concurrency(支持设置为 0
if input.Concurrency != nil {
account.Concurrency = *input.Concurrency
account.Concurrency = normalizeAccountConcurrency(account.Platform, account.Type, *input.Concurrency)
}
// 只在指针非 nil 时更新 Priority(支持设置为 0)
if input.Priority != nil {
@@ -108,6 +108,7 @@ func (s *GrokOAuthService) ExchangeCode(ctx context.Context, input *GrokExchange
if !ok {
return nil, infraerrors.New(http.StatusBadRequest, "GROK_OAUTH_SESSION_NOT_FOUND", "session not found or expired")
}
defer s.sessionStore.Delete(input.SessionID)
parsed := xai.ParseAuthorizationInput(input.Code)
code := strings.TrimSpace(parsed.Code)
@@ -118,6 +119,9 @@ func (s *GrokOAuthService) ExchangeCode(ctx context.Context, input *GrokExchange
if state == "" {
state = strings.TrimSpace(parsed.State)
}
if parsed.RequiresState && state == "" {
return nil, infraerrors.New(http.StatusBadRequest, "GROK_OAUTH_STATE_REQUIRED", "oauth state is required for callback URLs")
}
if state != "" && subtle.ConstantTimeCompare([]byte(state), []byte(session.State)) != 1 {
return nil, infraerrors.New(http.StatusBadRequest, "GROK_OAUTH_INVALID_STATE", "invalid oauth state")
}
@@ -139,7 +143,6 @@ func (s *GrokOAuthService) ExchangeCode(ctx context.Context, input *GrokExchange
if err != nil {
return nil, err
}
s.sessionStore.Delete(input.SessionID)
return s.tokenInfoFromResponse(tokenResp, session.ClientID, nil), nil
}
@@ -12,9 +12,11 @@ import (
type grokOAuthClientStub struct {
refreshResponse *xai.TokenResponse
exchangeCalls int
}
func (s *grokOAuthClientStub) ExchangeCode(context.Context, string, string, string, string, string) (*xai.TokenResponse, error) {
s.exchangeCalls++
return &xai.TokenResponse{}, nil
}
@@ -38,3 +40,29 @@ func TestGrokOAuthServiceRefreshTokenPreservesOriginalRefreshTokenWhenNotRotated
require.Equal(t, "original-refresh-token", info.RefreshToken)
require.Equal(t, "client-id", info.ClientID)
}
func TestGrokOAuthServiceExchangeCodeRequiresStateForCallbackURLAndConsumesSession(t *testing.T) {
client := &grokOAuthClientStub{}
svc := NewGrokOAuthService(nil, client)
defer svc.Stop()
auth, err := svc.GenerateAuthURL(context.Background(), nil, "")
require.NoError(t, err)
_, err = svc.ExchangeCode(context.Background(), &GrokExchangeCodeInput{
SessionID: auth.SessionID,
Code: "http://127.0.0.1:56121/callback?code=code-without-state",
})
require.Error(t, err)
require.Contains(t, err.Error(), "GROK_OAUTH_STATE_REQUIRED")
require.Zero(t, client.exchangeCalls)
_, err = svc.ExchangeCode(context.Background(), &GrokExchangeCodeInput{
SessionID: auth.SessionID,
Code: "code-with-state",
State: auth.State,
})
require.Error(t, err)
require.Contains(t, err.Error(), "GROK_OAUTH_SESSION_NOT_FOUND")
require.Zero(t, client.exchangeCalls)
}
@@ -132,7 +132,17 @@ func (p *GrokTokenProvider) markTempUnschedulable(account *Account, refreshErr e
}
now := time.Now()
until := now.Add(tokenRefreshTempUnschedDuration)
reason := "grok token refresh failed on request path: " + logredact.RedactText(refreshErr.Error())
redactedErr := "unknown error"
if refreshErr != nil {
redactedErr = logredact.RedactText(refreshErr.Error())
}
if isNonRetryableRefreshError(refreshErr) {
if err := p.accountRepo.SetError(context.Background(), account.ID, "grok token refresh failed (non-retryable): "+redactedErr); err != nil {
slog.Warn(grokTokenProviderLogComponent+".set_error_status_failed", "account_id", account.ID, "error", err)
}
return
}
reason := "grok token refresh failed on request path: " + redactedErr
bgCtx := context.Background()
if err := p.accountRepo.SetTempUnschedulable(bgCtx, account.ID, until, reason); err != nil {
slog.Warn(grokTokenProviderLogComponent+".set_temp_unschedulable_failed", "account_id", account.ID, "error", err)
@@ -27,6 +27,10 @@ func isOpenAIOAuthAccount(account *Account) bool {
return account != nil && account.Platform == PlatformOpenAI && account.Type == AccountTypeOAuth
}
func isGrokOAuthAccount(account *Account) bool {
return account != nil && account.Platform == PlatformGrok && account.Type == AccountTypeOAuth
}
func isOpenAIAccount(account *Account) bool {
return account != nil && (account.Platform == PlatformOpenAI || account.Platform == PlatformGrok)
}
@@ -172,6 +176,9 @@ func (s *OpenAIGatewayService) ShouldStopOpenAIOAuth429Failover(account *Account
if statusCode != http.StatusTooManyRequests || failedSwitches < openAIOAuth429StormMaxAccountSwitches {
return false
}
if isGrokOAuthAccount(account) {
return true
}
if !isOpenAIOAuthAccount(account) {
return false
}
@@ -121,3 +121,14 @@ func TestShouldStopOpenAIOAuth429Failover_OnlyDuringStorm(t *testing.T) {
require.False(t, svc.ShouldStopOpenAIOAuth429Failover(account, http.StatusInternalServerError, 1))
require.False(t, svc.ShouldStopOpenAIOAuth429Failover(account, http.StatusTooManyRequests, 0))
}
func TestShouldStopOpenAIOAuth429Failover_StopsGrokAfterFirst429Switch(t *testing.T) {
svc := &OpenAIGatewayService{}
account := &Account{ID: 44, Platform: PlatformGrok, Type: AccountTypeOAuth}
apiKeyAccount := &Account{ID: 45, Platform: PlatformGrok, Type: AccountTypeAPIKey}
require.True(t, svc.ShouldStopOpenAIOAuth429Failover(account, http.StatusTooManyRequests, 1))
require.False(t, svc.ShouldStopOpenAIOAuth429Failover(account, http.StatusTooManyRequests, 0))
require.False(t, svc.ShouldStopOpenAIOAuth429Failover(apiKeyAccount, http.StatusTooManyRequests, 1))
require.False(t, svc.ShouldStopOpenAIOAuth429Failover(account, http.StatusInternalServerError, 1))
}
@@ -10,6 +10,7 @@ import (
"time"
"github.com/Wei-Shaw/sub2api/internal/config"
"github.com/Wei-Shaw/sub2api/internal/util/logredact"
)
// tokenRefreshTempUnschedDuration token 刷新重试耗尽后临时不可调度的持续时间
@@ -309,7 +310,7 @@ func (s *TokenRefreshService) refreshWithRetry(ctx context.Context, account *Acc
// 不可重试错误(invalid_grant/invalid_client 等)直接标记 error 状态并返回
if isNonRetryableRefreshError(err) {
errorMsg := fmt.Sprintf("Token refresh failed (non-retryable): %v", err)
errorMsg := "Token refresh failed (non-retryable): " + logredact.RedactText(err.Error())
s.notifyAccountSchedulingBlocked(account, time.Time{}, "token_refresh_non_retryable")
if setErr := s.accountRepo.SetError(ctx, account.ID, errorMsg); setErr != nil {
slog.Error("token_refresh.set_error_status_failed",
@@ -346,7 +347,10 @@ func (s *TokenRefreshService) refreshWithRetry(ctx context.Context, account *Acc
// 设置临时不可调度 10 分钟(不标记 error,保持 status=active 让下个刷新周期能继续尝试)
until := time.Now().Add(tokenRefreshTempUnschedDuration)
reason := fmt.Sprintf("token refresh retry exhausted: %v", lastErr)
reason := "token refresh retry exhausted"
if lastErr != nil {
reason += ": " + logredact.RedactText(lastErr.Error())
}
s.notifyAccountSchedulingBlocked(account, until, "token_refresh_retry_exhausted")
if setErr := s.accountRepo.SetTempUnschedulable(ctx, account.ID, until, reason); setErr != nil {
slog.Warn("token_refresh.set_temp_unschedulable_failed",
@@ -450,6 +454,12 @@ func isNonRetryableRefreshError(err error) bool {
"access_denied", // 访问被拒绝
"missing_project_id", // 缺少 project_id
"no refresh token available",
"grok_oauth_entitlement_denied",
"entitlement_denied",
"invalid_scope",
"unknown scope",
"subscription required",
"no active grok subscription",
}
for _, needle := range nonRetryable {
if strings.Contains(msg, needle) {
@@ -538,6 +538,8 @@ func TestIsNonRetryableRefreshError(t *testing.T) {
{name: "unauthorized_client", err: errors.New("unauthorized_client"), expected: true},
{name: "access_denied", err: errors.New("access_denied"), expected: true},
{name: "no_refresh_token", err: errors.New("no refresh token available"), expected: true},
{name: "grok_entitlement_denied", err: errors.New("GROK_OAUTH_ENTITLEMENT_DENIED: subscription required"), expected: true},
{name: "invalid_scope", err: errors.New("invalid_scope: requested scope is not allowed"), expected: true},
{name: "invalid_grant_with_desc", err: errors.New("Error: invalid_grant - token revoked"), expected: true},
{name: "case_insensitive", err: errors.New("INVALID_GRANT"), expected: true},
}
@@ -3901,6 +3901,8 @@ watch(
accountCategory.value = 'oauth-based'
addMethod.value = 'oauth'
modelRestrictionMode.value = 'mapping'
form.concurrency = 1
form.load_factor = null
}
if (newPlatform !== 'gemini' && newPlatform !== 'anthropic' && accountCategory.value === 'service_account') {
accountCategory.value = 'oauth-based'