From 44f502bab8b2ac26ec3c51bccb94dee371bd41c2 Mon Sep 17 00:00:00 2001 From: Heatherm Huang Date: Fri, 26 Jun 2026 14:26:43 +0800 Subject: [PATCH] fix: address grok review feedback --- backend/cmd/server/wire_gen.go | 2 +- .../handler/admin/grok_oauth_handler_test.go | 4 +- .../internal/service/account_usage_service.go | 3 +- .../service/admin_account_concurrency_test.go | 27 ++++++++ backend/internal/service/admin_service.go | 9 ++- .../internal/service/grok_quota_service.go | 22 +++++-- .../service/grok_quota_service_test.go | 62 +++++++++++++++++-- .../internal/service/grok_token_provider.go | 9 +-- .../service/grok_token_provider_test.go | 4 +- .../internal/service/grok_token_refresher.go | 4 +- backend/internal/service/oauth_service.go | 6 ++ .../service/openai_gateway_grok_test.go | 6 +- .../service/openai_oauth_passthrough_test.go | 10 +-- backend/internal/service/wire.go | 5 +- 14 files changed, 140 insertions(+), 33 deletions(-) create mode 100644 backend/internal/service/admin_account_concurrency_test.go diff --git a/backend/cmd/server/wire_gen.go b/backend/cmd/server/wire_gen.go index 893a1d763a..ee3ccf9489 100644 --- a/backend/cmd/server/wire_gen.go +++ b/backend/cmd/server/wire_gen.go @@ -199,7 +199,7 @@ func initializeApplication(buildInfo handler.BuildInfo) (*Application, error) { openAIOAuthHandler := admin.NewOpenAIOAuthHandler(openAIOAuthService, adminService, openAIQuotaService) geminiOAuthHandler := admin.NewGeminiOAuthHandler(geminiOAuthService) antigravityOAuthHandler := admin.NewAntigravityOAuthHandler(antigravityOAuthService) - grokQuotaService := service.ProvideGrokQuotaService(accountRepository, grokTokenProvider, httpUpstream) + grokQuotaService := service.ProvideGrokQuotaService(accountRepository, proxyRepository, grokTokenProvider, httpUpstream) grokOAuthHandler := admin.NewGrokOAuthHandler(grokOAuthService, adminService, grokQuotaService) proxyHandler := admin.NewProxyHandler(adminService) adminRedeemHandler := admin.NewRedeemHandler(adminService, redeemService) diff --git a/backend/internal/handler/admin/grok_oauth_handler_test.go b/backend/internal/handler/admin/grok_oauth_handler_test.go index 3f2e7c23b5..6ac77e0e56 100644 --- a/backend/internal/handler/admin/grok_oauth_handler_test.go +++ b/backend/internal/handler/admin/grok_oauth_handler_test.go @@ -85,7 +85,7 @@ func TestGrokOAuthHandlerQueryQuotaProbesUpstream(t *testing.T) { }, Body: io.NopCloser(strings.NewReader(`{"id":"resp_probe"}`)), }} - quotaService := service.NewGrokQuotaService(repo, service.NewGrokTokenProvider(repo, nil, nil), upstream) + quotaService := service.NewGrokQuotaService(repo, nil, service.NewGrokTokenProvider(repo, nil), upstream) handler := NewGrokOAuthHandler(nil, nil, quotaService) router := gin.New() @@ -112,7 +112,7 @@ func TestGrokOAuthHandlerResetQuotaReturnsUnsupported(t *testing.T) { Platform: service.PlatformGrok, Type: service.AccountTypeOAuth, }} - quotaService := service.NewGrokQuotaService(repo, nil, nil) + quotaService := service.NewGrokQuotaService(repo, nil, nil, nil) handler := NewGrokOAuthHandler(nil, nil, quotaService) router := gin.New() diff --git a/backend/internal/service/account_usage_service.go b/backend/internal/service/account_usage_service.go index e47abe07f3..0087958ec5 100644 --- a/backend/internal/service/account_usage_service.go +++ b/backend/internal/service/account_usage_service.go @@ -864,7 +864,8 @@ func (s *AccountUsageService) getAntigravityUsage(ctx context.Context, account * func (s *AccountUsageService) getGrokUsage(ctx context.Context, account *Account) (*UsageInfo, error) { if s.grokQuotaFetcher == nil { - s.grokQuotaFetcher = NewGrokQuotaFetcher() + now := time.Now() + return &UsageInfo{UpdatedAt: &now}, nil } usage := s.grokQuotaFetcher.BuildUsageInfo(account) if usage.GrokQuotaSnapshotState == "" { diff --git a/backend/internal/service/admin_account_concurrency_test.go b/backend/internal/service/admin_account_concurrency_test.go new file mode 100644 index 0000000000..3544f80e24 --- /dev/null +++ b/backend/internal/service/admin_account_concurrency_test.go @@ -0,0 +1,27 @@ +//go:build unit + +package service + +import ( + "testing" + + "github.com/Wei-Shaw/sub2api/internal/pkg/xai" + "github.com/stretchr/testify/require" +) + +func TestNormalizeAccountConcurrencyCapsGrokOAuthUnlessUnsafe(t *testing.T) { + t.Setenv(xai.EnvUnsafeAllowHighConcurrency, "") + + require.Equal(t, 1, normalizeAccountConcurrency(PlatformGrok, AccountTypeOAuth, 0)) + require.Equal(t, 1, normalizeAccountConcurrency(PlatformGrok, AccountTypeOAuth, -5)) + require.Equal(t, 1, normalizeAccountConcurrency(PlatformGrok, AccountTypeOAuth, 50)) + require.Equal(t, 2, normalizeAccountConcurrency(PlatformOpenAI, AccountTypeOAuth, 2)) + require.Equal(t, 2, normalizeAccountConcurrency(PlatformGrok, AccountTypeAPIKey, 2)) +} + +func TestNormalizeAccountConcurrencyAllowsGrokOAuthUnsafeOverride(t *testing.T) { + t.Setenv(xai.EnvUnsafeAllowHighConcurrency, "true") + + require.Equal(t, 50, normalizeAccountConcurrency(PlatformGrok, AccountTypeOAuth, 50)) + require.Equal(t, 1, normalizeAccountConcurrency(PlatformGrok, AccountTypeOAuth, 0)) +} diff --git a/backend/internal/service/admin_service.go b/backend/internal/service/admin_service.go index 8af3b44904..7a5637dc75 100644 --- a/backend/internal/service/admin_service.go +++ b/backend/internal/service/admin_service.go @@ -2573,8 +2573,13 @@ func (s *adminServiceImpl) GetAccountsByIDs(ctx context.Context, ids []int64) ([ } func normalizeAccountConcurrency(platform, accountType string, concurrency int) int { - if platform == PlatformGrok && accountType == AccountTypeOAuth && concurrency <= 0 { - return 1 + if platform == PlatformGrok && accountType == AccountTypeOAuth { + if concurrency <= 0 { + return 1 + } + if concurrency > 1 && !xai.AllowUnsafeHighConcurrency() { + return 1 + } } return concurrency } diff --git a/backend/internal/service/grok_quota_service.go b/backend/internal/service/grok_quota_service.go index ce45a65dcd..19b1a01a0a 100644 --- a/backend/internal/service/grok_quota_service.go +++ b/backend/internal/service/grok_quota_service.go @@ -37,17 +37,20 @@ type GrokQuotaResetResult struct { type GrokQuotaService struct { accountRepo AccountRepository + proxyRepo ProxyRepository tokenProvider *GrokTokenProvider httpUpstream HTTPUpstream } func NewGrokQuotaService( accountRepo AccountRepository, + proxyRepo ProxyRepository, tokenProvider *GrokTokenProvider, httpUpstream HTTPUpstream, ) *GrokQuotaService { return &GrokQuotaService{ accountRepo: accountRepo, + proxyRepo: proxyRepo, tokenProvider: tokenProvider, httpUpstream: httpUpstream, } @@ -134,11 +137,22 @@ func (s *GrokQuotaService) prepareProbe(ctx context.Context, accountID int64) (* return nil, "", "", infraerrors.New(http.StatusBadGateway, "GROK_QUOTA_TOKEN_UNAVAILABLE", "access token is empty") } - proxyURL := "" - if account.ProxyID != nil && account.Proxy != nil { - proxyURL = account.Proxy.URL() + return account, token, s.resolveProxyURL(ctx, account), nil +} + +func (s *GrokQuotaService) resolveProxyURL(ctx context.Context, account *Account) string { + if account == nil || account.ProxyID == nil { + return "" } - return account, token, proxyURL, nil + switch { + case account.Proxy != nil: + return account.Proxy.URL() + case s != nil && s.proxyRepo != nil: + if proxy, err := s.proxyRepo.GetByID(ctx, *account.ProxyID); err == nil && proxy != nil { + return proxy.URL() + } + } + return "" } func (s *GrokQuotaService) loadGrokOAuthAccount(ctx context.Context, accountID int64) (*Account, error) { diff --git a/backend/internal/service/grok_quota_service_test.go b/backend/internal/service/grok_quota_service_test.go index fbc49067fa..d1da2e50a5 100644 --- a/backend/internal/service/grok_quota_service_test.go +++ b/backend/internal/service/grok_quota_service_test.go @@ -40,6 +40,17 @@ func (r *grokQuotaAccountRepo) SetTempUnschedulable(_ context.Context, id int64, return nil } +type grokQuotaProxyRepo struct { + proxyRepoStub + proxies map[int64]*Proxy + calls int +} + +func (r *grokQuotaProxyRepo) GetByID(_ context.Context, id int64) (*Proxy, error) { + r.calls++ + return r.proxies[id], nil +} + func TestGrokQuotaServiceProbeUsageStoresHeaders(t *testing.T) { t.Parallel() @@ -69,7 +80,7 @@ func TestGrokQuotaServiceProbeUsageStoresHeaders(t *testing.T) { }, Body: io.NopCloser(strings.NewReader(`{"id":"resp_probe"}`)), }} - svc := NewGrokQuotaService(repo, NewGrokTokenProvider(repo, nil, nil), upstream) + svc := NewGrokQuotaService(repo, nil, NewGrokTokenProvider(repo, nil), upstream) result, err := svc.ProbeUsage(context.Background(), 42) require.NoError(t, err) @@ -90,6 +101,49 @@ func TestGrokQuotaServiceProbeUsageStoresHeaders(t *testing.T) { require.NotNil(t, repo.updates[42][grokQuotaSnapshotExtraKey]) } +func TestGrokQuotaServiceProbeUsageLoadsProxyWhenAccountEdgeMissing(t *testing.T) { + t.Parallel() + + proxyID := int64(7) + account := &Account{ + ID: 46, + Platform: PlatformGrok, + Type: AccountTypeOAuth, + Concurrency: 1, + ProxyID: &proxyID, + Credentials: map[string]any{ + "access_token": "access-token", + "expires_at": time.Now().Add(time.Hour).UTC().Format(time.RFC3339), + }, + } + repo := &grokQuotaAccountRepo{ + mockAccountRepoForPlatform: &mockAccountRepoForPlatform{ + accountsByID: map[int64]*Account{46: account}, + }, + } + proxyRepo := &grokQuotaProxyRepo{ + proxies: map[int64]*Proxy{ + proxyID: { + ID: proxyID, + Protocol: "http", + Host: "proxy.test", + Port: 3128, + }, + }, + } + upstream := &httpUpstreamRecorder{resp: &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{}, + Body: io.NopCloser(strings.NewReader(`{"id":"resp_probe"}`)), + }} + svc := NewGrokQuotaService(repo, proxyRepo, NewGrokTokenProvider(repo, nil), upstream) + + _, err := svc.ProbeUsage(context.Background(), 46) + require.NoError(t, err) + require.Equal(t, 1, proxyRepo.calls) + require.Equal(t, "http://proxy.test:3128", upstream.lastProxyURL) +} + func TestGrokQuotaServiceProbeUsageStoresNoHeadersState(t *testing.T) { t.Parallel() @@ -113,7 +167,7 @@ func TestGrokQuotaServiceProbeUsageStoresNoHeadersState(t *testing.T) { Header: http.Header{}, Body: io.NopCloser(strings.NewReader(`{"id":"resp_probe"}`)), }} - svc := NewGrokQuotaService(repo, NewGrokTokenProvider(repo, nil, nil), upstream) + svc := NewGrokQuotaService(repo, nil, NewGrokTokenProvider(repo, nil), upstream) result, err := svc.ProbeUsage(context.Background(), 45) require.NoError(t, err) @@ -153,7 +207,7 @@ func TestGrokQuotaServiceProbeUsageReturnsRateLimitedSnapshot(t *testing.T) { Header: http.Header{"Retry-After": []string{"45"}}, Body: io.NopCloser(strings.NewReader(`{"error":{"message":"rate limited"}}`)), }} - svc := NewGrokQuotaService(repo, NewGrokTokenProvider(repo, nil, nil), upstream) + svc := NewGrokQuotaService(repo, nil, NewGrokTokenProvider(repo, nil), upstream) result, err := svc.ProbeUsage(context.Background(), 43) require.NoError(t, err) @@ -176,7 +230,7 @@ func TestGrokQuotaServiceResetQuotaUnsupported(t *testing.T) { accountsByID: map[int64]*Account{44: account}, }, } - svc := NewGrokQuotaService(repo, nil, nil) + svc := NewGrokQuotaService(repo, nil, nil, nil) _, err := svc.ResetQuota(context.Background(), 44) require.Error(t, err) diff --git a/backend/internal/service/grok_token_provider.go b/backend/internal/service/grok_token_provider.go index b12aa820ed..f19ee42b88 100644 --- a/backend/internal/service/grok_token_provider.go +++ b/backend/internal/service/grok_token_provider.go @@ -23,7 +23,6 @@ type GrokTokenCache = GeminiTokenCache type GrokTokenProvider struct { accountRepo AccountRepository tokenCache GrokTokenCache - grokOAuthService *GrokOAuthService refreshAPI *OAuthRefreshAPI executor OAuthRefreshExecutor refreshPolicy ProviderRefreshPolicy @@ -33,13 +32,11 @@ type GrokTokenProvider struct { func NewGrokTokenProvider( accountRepo AccountRepository, tokenCache GrokTokenCache, - grokOAuthService *GrokOAuthService, ) *GrokTokenProvider { return &GrokTokenProvider{ - accountRepo: accountRepo, - tokenCache: tokenCache, - grokOAuthService: grokOAuthService, - refreshPolicy: AntigravityProviderRefreshPolicy(), + accountRepo: accountRepo, + tokenCache: tokenCache, + refreshPolicy: AntigravityProviderRefreshPolicy(), } } diff --git a/backend/internal/service/grok_token_provider_test.go b/backend/internal/service/grok_token_provider_test.go index b647fb908b..1f7e2d8e0e 100644 --- a/backend/internal/service/grok_token_provider_test.go +++ b/backend/internal/service/grok_token_provider_test.go @@ -76,7 +76,7 @@ func TestGrokTokenProviderRefreshesExpiredTokenOnRequestPath(t *testing.T) { }) defer oauthSvc.Stop() - provider := NewGrokTokenProvider(repo, cache, oauthSvc) + provider := NewGrokTokenProvider(repo, cache) provider.SetRefreshAPI(NewOAuthRefreshAPI(repo, cache), NewGrokTokenRefresher(oauthSvc)) token, err := provider.GetAccessToken(context.Background(), account) @@ -109,7 +109,7 @@ func TestGrokTokenProviderRefreshFailureUnschedulesWithRedactedReason(t *testing repo.accountsByID = map[int64]*Account{55: account} cache := &grokTokenCacheForProviderTest{lockResult: true} tempCache := &tempUnschedCacheStub{} - provider := NewGrokTokenProvider(repo, cache, nil) + provider := NewGrokTokenProvider(repo, cache) provider.SetRefreshAPI(NewOAuthRefreshAPI(repo, cache), &tokenRefresherStub{ err: errors.New("temporary refresh failure access_token=leaked-access refresh_token=leaked-refresh"), }) diff --git a/backend/internal/service/grok_token_refresher.go b/backend/internal/service/grok_token_refresher.go index 018667f200..92d88cc058 100644 --- a/backend/internal/service/grok_token_refresher.go +++ b/backend/internal/service/grok_token_refresher.go @@ -10,10 +10,10 @@ import ( const grokTokenRefreshSkew = time.Hour type GrokTokenRefresher struct { - grokOAuthService *GrokOAuthService + grokOAuthService GrokOAuthTokenService } -func NewGrokTokenRefresher(grokOAuthService *GrokOAuthService) *GrokTokenRefresher { +func NewGrokTokenRefresher(grokOAuthService GrokOAuthTokenService) *GrokTokenRefresher { return &GrokTokenRefresher{grokOAuthService: grokOAuthService} } diff --git a/backend/internal/service/oauth_service.go b/backend/internal/service/oauth_service.go index c03ba912dc..1369dd9e89 100644 --- a/backend/internal/service/oauth_service.go +++ b/backend/internal/service/oauth_service.go @@ -24,6 +24,12 @@ type GrokOAuthClient interface { RefreshToken(ctx context.Context, refreshToken, proxyURL, clientID string) (*xai.TokenResponse, error) } +// GrokOAuthTokenService is the narrow refresh port used by Grok token providers. +type GrokOAuthTokenService interface { + RefreshAccountToken(ctx context.Context, account *Account) (*GrokTokenInfo, error) + BuildAccountCredentials(tokenInfo *GrokTokenInfo) map[string]any +} + // ClaudeOAuthClient handles HTTP requests for Claude OAuth flows type ClaudeOAuthClient interface { GetOrganizationUUID(ctx context.Context, sessionKey, proxyURL string) (string, error) diff --git a/backend/internal/service/openai_gateway_grok_test.go b/backend/internal/service/openai_gateway_grok_test.go index 65b17586aa..15dd6a3acb 100644 --- a/backend/internal/service/openai_gateway_grok_test.go +++ b/backend/internal/service/openai_gateway_grok_test.go @@ -118,7 +118,7 @@ func TestForwardAsChatCompletionsForGrokUsesXAIChatCompletionsAndSnapshots(t *te }} svc := &OpenAIGatewayService{ httpUpstream: upstream, - grokTokenProvider: NewGrokTokenProvider(repo, nil, nil), + grokTokenProvider: NewGrokTokenProvider(repo, nil), accountRepo: repo, } @@ -182,7 +182,7 @@ func TestForwardGrokResponsesStreamingUsesXAIResponsesAndSnapshots(t *testing.T) }} svc := &OpenAIGatewayService{ httpUpstream: upstream, - grokTokenProvider: NewGrokTokenProvider(repo, nil, nil), + grokTokenProvider: NewGrokTokenProvider(repo, nil), accountRepo: repo, } @@ -251,7 +251,7 @@ func TestForwardAsChatCompletionsForGrokStreamingUsesRawXAIChatCompletions(t *te svc := &OpenAIGatewayService{ cfg: rawChatCompletionsTestConfig(), httpUpstream: upstream, - grokTokenProvider: NewGrokTokenProvider(repo, nil, nil), + grokTokenProvider: NewGrokTokenProvider(repo, nil), accountRepo: repo, } diff --git a/backend/internal/service/openai_oauth_passthrough_test.go b/backend/internal/service/openai_oauth_passthrough_test.go index b371808066..0b6c16130b 100644 --- a/backend/internal/service/openai_oauth_passthrough_test.go +++ b/backend/internal/service/openai_oauth_passthrough_test.go @@ -24,10 +24,11 @@ import ( func f64p(v float64) *float64 { return &v } type httpUpstreamRecorder struct { - lastReq *http.Request - lastBody []byte - requests []*http.Request - bodies [][]byte + lastReq *http.Request + lastBody []byte + lastProxyURL string + requests []*http.Request + bodies [][]byte resp *http.Response responses []*http.Response @@ -36,6 +37,7 @@ type httpUpstreamRecorder struct { func (u *httpUpstreamRecorder) Do(req *http.Request, proxyURL string, accountID int64, accountConcurrency int) (*http.Response, error) { u.lastReq = req + u.lastProxyURL = proxyURL if req != nil && req.Body != nil { b, _ := io.ReadAll(req.Body) u.lastBody = b diff --git a/backend/internal/service/wire.go b/backend/internal/service/wire.go index 43c0b887f5..b3b1170d97 100644 --- a/backend/internal/service/wire.go +++ b/backend/internal/service/wire.go @@ -127,10 +127,11 @@ func ProvideOpenAIQuotaService( func ProvideGrokQuotaService( accountRepo AccountRepository, + proxyRepo ProxyRepository, tokenProvider *GrokTokenProvider, httpUpstream HTTPUpstream, ) *GrokQuotaService { - return NewGrokQuotaService(accountRepo, tokenProvider, httpUpstream) + return NewGrokQuotaService(accountRepo, proxyRepo, tokenProvider, httpUpstream) } // ProvideGeminiTokenProvider creates GeminiTokenProvider with OAuthRefreshAPI injection @@ -171,7 +172,7 @@ func ProvideGrokTokenProvider( refreshAPI *OAuthRefreshAPI, tempUnschedCache TempUnschedCache, ) *GrokTokenProvider { - p := NewGrokTokenProvider(accountRepo, tokenCache, grokOAuthService) + p := NewGrokTokenProvider(accountRepo, tokenCache) executor := NewGrokTokenRefresher(grokOAuthService) p.SetRefreshAPI(refreshAPI, executor) p.SetRefreshPolicy(AntigravityProviderRefreshPolicy())