diff --git a/README.md b/README.md index db1c4f6c8a..84aab47ad5 100644 --- a/README.md +++ b/README.md @@ -581,6 +581,50 @@ Simple Mode is designed for individual developers or internal teams who want qui --- +## Grok / xAI OAuth Support + +Sub2API supports Grok subscription accounts through xAI OAuth and forwards OpenAI-compatible Responses traffic to xAI. + +### Supported Scope + +- Platform name: `grok` +- Account type: OAuth subscription accounts +- Public gateway target: `/v1/responses` and `/responses`, forwarded to `${XAI_BASE_URL:-https://api.x.ai/v1}/responses` +- Initial models: `grok-4.3`, `grok-build-0.1`, `grok-4.20-0309-reasoning`, `grok-4.20-0309-non-reasoning`, and `grok-4.20-multi-agent-0309` +- Out of scope for this provider: public Grok Chat Completions routes, image, video, TTS, transcription, browser automation, cookies, and Grok web scraping + +### OAuth Configuration + +The Grok OAuth flow uses PKCE and does not require committing private secrets. The default client details follow the public xAI OAuth flow used by compatible clients, and every value can be overridden by environment variable: + +| Variable | Default | +|----------|---------| +| `XAI_OAUTH_CLIENT_ID` | Public xAI OAuth client ID | +| `XAI_OAUTH_SCOPE` | `openid profile email offline_access grok-cli:access api:access` | +| `XAI_OAUTH_REDIRECT_URI` | `http://127.0.0.1:56121/callback` | +| `XAI_OAUTH_AUTHORIZE_URL` | `https://auth.x.ai/oauth2/authorize` | +| `XAI_OAUTH_TOKEN_URL` | `https://auth.x.ai/oauth2/token` | +| `XAI_BASE_URL` | `https://api.x.ai/v1` | + +Administrators can create or reauthorize Grok accounts from the dashboard, or use the admin API: + +| Endpoint | Purpose | +|----------|---------| +| `POST /api/v1/admin/grok/oauth/auth-url` | Generate an xAI OAuth authorization URL | +| `POST /api/v1/admin/grok/oauth/exchange-code` | Exchange a callback URL, query string, or code for OAuth credentials | +| `POST /api/v1/admin/grok/oauth/refresh-token` | Validate or refresh a Grok refresh token | +| `POST /api/v1/admin/grok/accounts/:id/refresh` | Refresh an existing Grok account | + +Credential storage reuses the existing account JSON fields: `access_token`, `refresh_token`, `token_type`, `expires_at`, optional `email`, optional `subscription_tier`, and `entitlement_status`. + +### Usage And Quota Display + +xAI quota is passive. Sub2API does not invent subscription quota values; it records whitelisted xAI rate-limit headers from successful or rate-limited upstream responses when xAI sends them. Before the first usable upstream response, the dashboard shows quota as unknown and still displays local Sub2API usage stats. + +`401` responses mark the account as needing reauthorization. `403` responses are treated as entitlement or subscription-tier failures instead of token-refresh loops. `429` responses use `Retry-After` or a short cooldown to temporarily remove the account from scheduling. + +--- + ## Antigravity Support Sub2API supports [Antigravity](https://antigravity.so/) accounts. After authorization, dedicated endpoints are available for Claude and Gemini models. diff --git a/backend/cmd/server/wire.go b/backend/cmd/server/wire.go index da119aaaca..b9a9a3e80e 100644 --- a/backend/cmd/server/wire.go +++ b/backend/cmd/server/wire.go @@ -94,6 +94,7 @@ func provideCleanup( openaiOAuth *service.OpenAIOAuthService, geminiOAuth *service.GeminiOAuthService, antigravityOAuth *service.AntigravityOAuthService, + grokOAuth *service.GrokOAuthService, openAIGateway *service.OpenAIGatewayService, scheduledTestRunner *service.ScheduledTestRunnerService, backupSvc *service.BackupService, @@ -222,6 +223,12 @@ func provideCleanup( antigravityOAuth.Stop() return nil }}, + {"GrokOAuthService", func() error { + if grokOAuth != nil { + grokOAuth.Stop() + } + return nil + }}, {"OpenAIWSPool", func() error { if openAIGateway != nil { openAIGateway.CloseOpenAIWSPool() diff --git a/backend/cmd/server/wire_gen.go b/backend/cmd/server/wire_gen.go index eefd228ef6..ee3ccf9489 100644 --- a/backend/cmd/server/wire_gen.go +++ b/backend/cmd/server/wire_gen.go @@ -141,7 +141,10 @@ func initializeApplication(buildInfo handler.BuildInfo) (*Application, error) { privacyClientFactory := providePrivacyClientFactory() openAIOAuthService := service.ProvideOpenAIOAuthService(proxyRepository, openAIOAuthClient, privacyClientFactory) openAITokenProvider := service.ProvideOpenAITokenProvider(accountRepository, geminiTokenCache, openAIOAuthService, oAuthRefreshAPI) - openAIGatewayService := service.NewOpenAIGatewayService(accountRepository, usageLogRepository, usageBillingRepository, userRepository, userSubscriptionRepository, userGroupRateRepository, gatewayCache, configConfig, schedulerSnapshotService, concurrencyService, billingService, rateLimitService, billingCacheService, httpUpstream, deferredService, openAITokenProvider, modelPricingResolver, channelService, balanceNotifyService, settingService, serviceUserPlatformQuotaRepository) + grokOAuthClient := repository.NewGrokOAuthClient() + grokOAuthService := service.NewGrokOAuthService(proxyRepository, grokOAuthClient) + grokTokenProvider := service.ProvideGrokTokenProvider(accountRepository, geminiTokenCache, grokOAuthService, oAuthRefreshAPI, tempUnschedCache) + openAIGatewayService := service.NewOpenAIGatewayService(accountRepository, usageLogRepository, usageBillingRepository, userRepository, userSubscriptionRepository, userGroupRateRepository, gatewayCache, configConfig, schedulerSnapshotService, concurrencyService, billingService, rateLimitService, billingCacheService, httpUpstream, deferredService, openAITokenProvider, grokTokenProvider, modelPricingResolver, channelService, balanceNotifyService, settingService, serviceUserPlatformQuotaRepository) geminiOAuthClient := repository.NewGeminiOAuthClient(configConfig) geminiCliCodeAssistClient := repository.NewGeminiCliCodeAssistClient() driveClient := repository.NewGeminiDriveClient() @@ -178,8 +181,9 @@ func initializeApplication(buildInfo handler.BuildInfo) (*Application, error) { groupHandler := admin.NewGroupHandler(adminService, dashboardService, groupCapacityService) claudeUsageFetcher := repository.NewClaudeUsageFetcher(httpUpstream) antigravityQuotaFetcher := service.NewAntigravityQuotaFetcher(proxyRepository) + grokQuotaFetcher := service.NewGrokQuotaFetcher() usageCache := service.NewUsageCache() - accountUsageService := service.NewAccountUsageService(accountRepository, usageLogRepository, claudeUsageFetcher, geminiQuotaService, antigravityQuotaFetcher, usageCache, identityCache, tlsFingerprintProfileService) + accountUsageService := service.NewAccountUsageService(accountRepository, usageLogRepository, claudeUsageFetcher, geminiQuotaService, antigravityQuotaFetcher, grokQuotaFetcher, usageCache, identityCache, tlsFingerprintProfileService) accountTestService := service.NewAccountTestService(accountRepository, geminiTokenProvider, claudeTokenProvider, antigravityGatewayService, httpUpstream, configConfig, tlsFingerprintProfileService) crsSyncService := service.NewCRSSyncService(accountRepository, proxyRepository, oAuthService, openAIOAuthService, geminiOAuthService, configConfig) accountHandler := admin.NewAccountHandler(adminService, oAuthService, openAIOAuthService, geminiOAuthService, antigravityOAuthService, rateLimitService, accountUsageService, accountTestService, concurrencyService, crsSyncService, sessionLimitCache, rpmCache, compositeTokenCacheInvalidator) @@ -195,6 +199,8 @@ 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, proxyRepository, grokTokenProvider, httpUpstream) + grokOAuthHandler := admin.NewGrokOAuthHandler(grokOAuthService, adminService, grokQuotaService) proxyHandler := admin.NewProxyHandler(adminService) adminRedeemHandler := admin.NewRedeemHandler(adminService, redeemService) promoHandler := admin.NewPromoHandler(promoService) @@ -242,7 +248,7 @@ func initializeApplication(buildInfo handler.BuildInfo) (*Application, error) { paymentHandler := admin.NewPaymentHandler(paymentService, paymentConfigService) affiliateHandler := admin.NewAffiliateHandler(affiliateService, adminService) complianceHandler := admin.NewComplianceHandler(settingService) - adminHandlers := handler.ProvideAdminHandlers(dashboardHandler, adminUserHandler, groupHandler, accountHandler, adminAnnouncementHandler, dataManagementHandler, backupHandler, oAuthHandler, openAIOAuthHandler, geminiOAuthHandler, antigravityOAuthHandler, proxyHandler, adminRedeemHandler, promoHandler, settingHandler, opsHandler, systemHandler, adminSubscriptionHandler, adminUsageHandler, userAttributeHandler, errorPassthroughHandler, tlsFingerprintProfileHandler, adminAPIKeyHandler, scheduledTestHandler, channelHandler, channelMonitorHandler, channelMonitorRequestTemplateHandler, contentModerationHandler, paymentHandler, affiliateHandler, complianceHandler) + adminHandlers := handler.ProvideAdminHandlers(dashboardHandler, adminUserHandler, groupHandler, accountHandler, adminAnnouncementHandler, dataManagementHandler, backupHandler, oAuthHandler, openAIOAuthHandler, geminiOAuthHandler, antigravityOAuthHandler, grokOAuthHandler, proxyHandler, adminRedeemHandler, promoHandler, settingHandler, opsHandler, systemHandler, adminSubscriptionHandler, adminUsageHandler, userAttributeHandler, errorPassthroughHandler, tlsFingerprintProfileHandler, adminAPIKeyHandler, scheduledTestHandler, channelHandler, channelMonitorHandler, channelMonitorRequestTemplateHandler, contentModerationHandler, paymentHandler, affiliateHandler, complianceHandler) usageRecordWorkerPool := service.NewUsageRecordWorkerPool(configConfig) userMsgQueueCache := repository.NewUserMsgQueueCache(redisClient) userMessageQueueService := service.ProvideUserMessageQueueService(userMsgQueueCache, rpmCache, configConfig) @@ -266,7 +272,7 @@ func initializeApplication(buildInfo handler.BuildInfo) (*Application, error) { opsAlertEvaluatorService := service.ProvideOpsAlertEvaluatorService(opsService, opsRepository, emailService, redisClient, configConfig, proxyRepository) opsCleanupService := service.ProvideOpsCleanupService(opsRepository, db, redisClient, configConfig, channelMonitorService, settingRepository, opsService) opsScheduledReportService := service.ProvideOpsScheduledReportService(opsService, userService, emailService, redisClient, configConfig) - tokenRefreshService := service.ProvideTokenRefreshService(accountRepository, oAuthService, openAIOAuthService, geminiOAuthService, antigravityOAuthService, compositeTokenCacheInvalidator, schedulerCache, configConfig, tempUnschedCache, privacyClientFactory, proxyRepository, oAuthRefreshAPI, openAIGatewayService) + tokenRefreshService := service.ProvideTokenRefreshService(accountRepository, oAuthService, openAIOAuthService, geminiOAuthService, antigravityOAuthService, grokOAuthService, compositeTokenCacheInvalidator, schedulerCache, configConfig, tempUnschedCache, privacyClientFactory, proxyRepository, oAuthRefreshAPI, openAIGatewayService) accountExpiryService := service.ProvideAccountExpiryService(accountRepository) proxyExpiryService := service.ProvideProxyExpiryService(proxyRepository) subscriptionExpiryService := service.ProvideSubscriptionExpiryService(userSubscriptionRepository, settingRepository, notificationEmailService, leaderLockCache, db) @@ -274,7 +280,7 @@ func initializeApplication(buildInfo handler.BuildInfo) (*Application, error) { paymentOrderExpiryService := service.ProvidePaymentOrderExpiryService(paymentService, leaderLockCache, db) channelMonitorRunner := service.ProvideChannelMonitorRunner(channelMonitorService, settingService) userPlatformQuotaUsageFlusher := service.ProvideUserPlatformQuotaUsageFlusher(configConfig, billingCache, serviceUserPlatformQuotaRepository, timingWheelService) - v := provideCleanup(client, redisClient, opsMetricsCollector, opsAggregationService, opsAlertEvaluatorService, opsCleanupService, opsScheduledReportService, opsSystemLogSink, schedulerSnapshotService, tokenRefreshService, accountExpiryService, proxyExpiryService, subscriptionExpiryService, usageCleanupService, idempotencyCleanupService, pricingService, emailQueueService, billingCacheService, usageRecordWorkerPool, subscriptionService, oAuthService, openAIOAuthService, geminiOAuthService, antigravityOAuthService, openAIGatewayService, scheduledTestRunnerService, backupService, paymentOrderExpiryService, channelMonitorRunner, userPlatformQuotaUsageFlusher) + v := provideCleanup(client, redisClient, opsMetricsCollector, opsAggregationService, opsAlertEvaluatorService, opsCleanupService, opsScheduledReportService, opsSystemLogSink, schedulerSnapshotService, tokenRefreshService, accountExpiryService, proxyExpiryService, subscriptionExpiryService, usageCleanupService, idempotencyCleanupService, pricingService, emailQueueService, billingCacheService, usageRecordWorkerPool, subscriptionService, oAuthService, openAIOAuthService, geminiOAuthService, antigravityOAuthService, grokOAuthService, openAIGatewayService, scheduledTestRunnerService, backupService, paymentOrderExpiryService, channelMonitorRunner, userPlatformQuotaUsageFlusher) application := &Application{ Server: httpServer, Cleanup: v, @@ -325,6 +331,7 @@ func provideCleanup( openaiOAuth *service.OpenAIOAuthService, geminiOAuth *service.GeminiOAuthService, antigravityOAuth *service.AntigravityOAuthService, + grokOAuth *service.GrokOAuthService, openAIGateway *service.OpenAIGatewayService, scheduledTestRunner *service.ScheduledTestRunnerService, backupSvc *service.BackupService, @@ -452,6 +459,12 @@ func provideCleanup( antigravityOAuth.Stop() return nil }}, + {"GrokOAuthService", func() error { + if grokOAuth != nil { + grokOAuth.Stop() + } + return nil + }}, {"OpenAIWSPool", func() error { if openAIGateway != nil { openAIGateway.CloseOpenAIWSPool() diff --git a/backend/cmd/server/wire_gen_test.go b/backend/cmd/server/wire_gen_test.go index 4a00c1cceb..ef74cb4a2d 100644 --- a/backend/cmd/server/wire_gen_test.go +++ b/backend/cmd/server/wire_gen_test.go @@ -74,6 +74,7 @@ func TestProvideCleanup_WithMinimalDependencies_NoPanic(t *testing.T) { openAIOAuthSvc, geminiOAuthSvc, antigravityOAuthSvc, + nil, // grokOAuth nil, // openAIGateway nil, // scheduledTestRunner nil, // backupSvc diff --git a/backend/ent/schema/user_platform_quota.go b/backend/ent/schema/user_platform_quota.go index 8fd8acc016..a0b5598600 100644 --- a/backend/ent/schema/user_platform_quota.go +++ b/backend/ent/schema/user_platform_quota.go @@ -41,7 +41,7 @@ func (UserPlatformQuota) Fields() []ent.Field { // 注意:平台列表的单一权威源为 service.AllowedQuotaPlatforms; // 此处为 ent 构建期约束,需与 service.AllowedQuotaPlatforms 保持同步。 switch s { - case "anthropic", "openai", "gemini", "antigravity": + case "anthropic", "openai", "gemini", "antigravity", "grok": return nil default: return fmt.Errorf("platform %q is not allowed", s) diff --git a/backend/go.sum b/backend/go.sum index 7735fda29e..fbc04494ce 100644 --- a/backend/go.sum +++ b/backend/go.sum @@ -164,6 +164,8 @@ github.com/google/go-querystring v1.1.0/go.mod h1:Kcdr2DB4koayq7X8pmAG4sNG59So17 github.com/google/gofuzz v1.0.0/go.mod h1:dBl0BpW6vV/+mYPU4Po3pmUjxk6FQPldtuIdl/M65Eg= github.com/google/pprof v0.0.0-20250317173921-a4b03ec1a45e h1:ijClszYn+mADRFY17kjQEVQ1XRhq2/JR1M3sGqeJoxs= github.com/google/pprof v0.0.0-20250317173921-a4b03ec1a45e/go.mod h1:boTsfXsheKC2y+lKOCMpSfarhxDeIzfZG1jqGcPl3cA= +github.com/google/subcommands v1.2.0 h1:vWQspBTo2nEqTUFita5/KeEWlUL8kQObDFbub/EN9oE= +github.com/google/subcommands v1.2.0/go.mod h1:ZjhPrFU+Olkh9WazFPsl27BQ4UPiG37m3yTrtFlrHVk= github.com/google/uuid v1.6.0 h1:NIvaJDMOsjHA8n1jAhLSgzrAzy1Hgr+hNrb57e+94F0= github.com/google/uuid v1.6.0/go.mod h1:TIyPZe4MgqvfeYDBFedMoGGpEw/LqOeaOT+nhxU+yHo= github.com/google/wire v0.7.0 h1:JxUKI6+CVBgCO2WToKy/nQk0sS+amI9z9EjVmdaocj4= diff --git a/backend/internal/domain/constants.go b/backend/internal/domain/constants.go index 4ed3f72b3a..844348c024 100644 --- a/backend/internal/domain/constants.go +++ b/backend/internal/domain/constants.go @@ -22,6 +22,7 @@ const ( PlatformOpenAI = "openai" PlatformGemini = "gemini" PlatformAntigravity = "antigravity" + PlatformGrok = "grok" ) // Account type constants diff --git a/backend/internal/handler/admin/channel_handler.go b/backend/internal/handler/admin/channel_handler.go index bf547346bd..30d208c0d2 100644 --- a/backend/internal/handler/admin/channel_handler.go +++ b/backend/internal/handler/admin/channel_handler.go @@ -509,6 +509,7 @@ var platformToLiteLLMProvider = map[string]string{ service.PlatformOpenAI: "openai", service.PlatformGemini: "google", service.PlatformAntigravity: "anthropic", + service.PlatformGrok: "xai", } // SyncPricingModels 返回 LiteLLM 定价目录中指定平台的最新模型列表 diff --git a/backend/internal/handler/admin/grok_oauth_handler.go b/backend/internal/handler/admin/grok_oauth_handler.go new file mode 100644 index 0000000000..dafa3076b8 --- /dev/null +++ b/backend/internal/handler/admin/grok_oauth_handler.go @@ -0,0 +1,246 @@ +package admin + +import ( + "strconv" + "strings" + + "github.com/Wei-Shaw/sub2api/internal/handler/dto" + "github.com/Wei-Shaw/sub2api/internal/pkg/response" + "github.com/Wei-Shaw/sub2api/internal/pkg/xai" + "github.com/Wei-Shaw/sub2api/internal/service" + "github.com/gin-gonic/gin" +) + +type GrokOAuthHandler struct { + grokOAuthService *service.GrokOAuthService + adminService service.AdminService + quotaService *service.GrokQuotaService +} + +func NewGrokOAuthHandler( + grokOAuthService *service.GrokOAuthService, + adminService service.AdminService, + quotaService *service.GrokQuotaService, +) *GrokOAuthHandler { + return &GrokOAuthHandler{ + grokOAuthService: grokOAuthService, + adminService: adminService, + quotaService: quotaService, + } +} + +type GrokGenerateAuthURLRequest struct { + ProxyID *int64 `json:"proxy_id"` + RedirectURI string `json:"redirect_uri"` +} + +func (h *GrokOAuthHandler) GenerateAuthURL(c *gin.Context) { + var req GrokGenerateAuthURLRequest + if err := c.ShouldBindJSON(&req); err != nil { + req = GrokGenerateAuthURLRequest{} + } + result, err := h.grokOAuthService.GenerateAuthURL(c.Request.Context(), req.ProxyID, req.RedirectURI) + if err != nil { + response.ErrorFrom(c, err) + return + } + response.Success(c, result) +} + +type GrokExchangeCodeRequest struct { + SessionID string `json:"session_id" binding:"required"` + Code string `json:"code" binding:"required"` + State string `json:"state"` + RedirectURI string `json:"redirect_uri"` + ProxyID *int64 `json:"proxy_id"` +} + +func (h *GrokOAuthHandler) ExchangeCode(c *gin.Context) { + var req GrokExchangeCodeRequest + if err := c.ShouldBindJSON(&req); err != nil { + response.BadRequest(c, "Invalid request: "+err.Error()) + return + } + tokenInfo, err := h.grokOAuthService.ExchangeCode(c.Request.Context(), &service.GrokExchangeCodeInput{ + SessionID: req.SessionID, + Code: req.Code, + State: req.State, + RedirectURI: req.RedirectURI, + ProxyID: req.ProxyID, + }) + if err != nil { + response.ErrorFrom(c, err) + return + } + response.Success(c, tokenInfo) +} + +type GrokRefreshTokenRequest struct { + RefreshToken string `json:"refresh_token"` + RT string `json:"rt"` + ClientID string `json:"client_id"` + ProxyID *int64 `json:"proxy_id"` +} + +func (h *GrokOAuthHandler) RefreshToken(c *gin.Context) { + var req GrokRefreshTokenRequest + if err := c.ShouldBindJSON(&req); err != nil { + response.BadRequest(c, "Invalid request: "+err.Error()) + return + } + refreshToken := strings.TrimSpace(req.RefreshToken) + if refreshToken == "" { + refreshToken = strings.TrimSpace(req.RT) + } + if refreshToken == "" { + response.BadRequest(c, "refresh_token is required") + return + } + + var proxyURL string + if req.ProxyID != nil { + proxy, err := h.adminService.GetProxy(c.Request.Context(), *req.ProxyID) + if err == nil && proxy != nil { + proxyURL = proxy.URL() + } + } + tokenInfo, err := h.grokOAuthService.RefreshToken(c.Request.Context(), refreshToken, proxyURL, req.ClientID) + if err != nil { + response.ErrorFrom(c, err) + return + } + response.Success(c, tokenInfo) +} + +func (h *GrokOAuthHandler) RefreshAccountToken(c *gin.Context) { + accountID, err := strconv.ParseInt(c.Param("id"), 10, 64) + if err != nil { + response.BadRequest(c, "Invalid account ID") + return + } + account, err := h.adminService.GetAccount(c.Request.Context(), accountID) + if err != nil { + response.ErrorFrom(c, err) + return + } + if account.Platform != service.PlatformGrok { + response.BadRequest(c, "Account platform does not match Grok OAuth endpoint") + return + } + if !account.IsOAuth() { + response.BadRequest(c, "Cannot refresh non-OAuth account credentials") + return + } + tokenInfo, err := h.grokOAuthService.RefreshAccountToken(c.Request.Context(), account) + if err != nil { + response.ErrorFrom(c, err) + return + } + newCredentials := h.grokOAuthService.BuildAccountCredentials(tokenInfo) + newCredentials = service.MergeCredentials(account.Credentials, newCredentials) + if baseURL := strings.TrimSpace(account.GetCredential("base_url")); baseURL != "" { + newCredentials["base_url"] = baseURL + } + updatedAccount, err := h.adminService.UpdateAccount(c.Request.Context(), accountID, &service.UpdateAccountInput{ + Credentials: newCredentials, + }) + if err != nil { + response.ErrorFrom(c, err) + return + } + response.Success(c, dto.AccountFromService(updatedAccount)) +} + +func (h *GrokOAuthHandler) CreateAccountFromOAuth(c *gin.Context) { + var req struct { + SessionID string `json:"session_id" binding:"required"` + Code string `json:"code" binding:"required"` + State string `json:"state"` + RedirectURI string `json:"redirect_uri"` + ProxyID *int64 `json:"proxy_id"` + Name string `json:"name"` + Concurrency int `json:"concurrency"` + Priority int `json:"priority"` + GroupIDs []int64 `json:"group_ids"` + } + if err := c.ShouldBindJSON(&req); err != nil { + response.BadRequest(c, "Invalid request: "+err.Error()) + return + } + tokenInfo, err := h.grokOAuthService.ExchangeCode(c.Request.Context(), &service.GrokExchangeCodeInput{ + SessionID: req.SessionID, + Code: req.Code, + State: req.State, + RedirectURI: req.RedirectURI, + ProxyID: req.ProxyID, + }) + if err != nil { + response.ErrorFrom(c, err) + return + } + credentials := h.grokOAuthService.BuildAccountCredentials(tokenInfo) + + name := strings.TrimSpace(req.Name) + if name == "" && tokenInfo.Email != "" { + name = tokenInfo.Email + } + if name == "" { + name = "Grok OAuth Account" + } + + account, err := h.adminService.CreateAccount(c.Request.Context(), &service.CreateAccountInput{ + Name: name, + Platform: service.PlatformGrok, + Type: service.AccountTypeOAuth, + Credentials: credentials, + ProxyID: req.ProxyID, + Concurrency: req.Concurrency, + Priority: req.Priority, + GroupIDs: req.GroupIDs, + }) + if err != nil { + response.ErrorFrom(c, err) + return + } + response.Success(c, dto.AccountFromService(account)) +} + +func (h *GrokOAuthHandler) QueryQuota(c *gin.Context) { + accountID, err := strconv.ParseInt(c.Param("id"), 10, 64) + if err != nil { + response.BadRequest(c, "Invalid account ID") + return + } + if h.quotaService == nil { + response.BadRequest(c, "grok quota service is not enabled") + return + } + result, err := h.quotaService.ProbeUsage(c.Request.Context(), accountID) + if err != nil { + response.ErrorFrom(c, err) + return + } + response.Success(c, result) +} + +func (h *GrokOAuthHandler) ResetQuota(c *gin.Context) { + accountID, err := strconv.ParseInt(c.Param("id"), 10, 64) + if err != nil { + response.BadRequest(c, "Invalid account ID") + return + } + if h.quotaService == nil { + response.BadRequest(c, "grok quota service is not enabled") + return + } + result, err := h.quotaService.ResetQuota(c.Request.Context(), accountID) + if err != nil { + response.ErrorFrom(c, err) + return + } + response.Success(c, result) +} + +func (h *GrokOAuthHandler) RuntimeSanity(c *gin.Context) { + response.Success(c, xai.RuntimeSanity()) +} diff --git a/backend/internal/handler/admin/grok_oauth_handler_test.go b/backend/internal/handler/admin/grok_oauth_handler_test.go new file mode 100644 index 0000000000..6ac77e0e56 --- /dev/null +++ b/backend/internal/handler/admin/grok_oauth_handler_test.go @@ -0,0 +1,147 @@ +//go:build unit + +package admin + +import ( + "context" + "io" + "net/http" + "net/http/httptest" + "strings" + "testing" + "time" + + "github.com/gin-gonic/gin" + "github.com/stretchr/testify/require" + + "github.com/Wei-Shaw/sub2api/internal/pkg/tlsfingerprint" + "github.com/Wei-Shaw/sub2api/internal/pkg/xai" + "github.com/Wei-Shaw/sub2api/internal/service" +) + +type grokQuotaHandlerAccountRepo struct { + service.AccountRepository + account *service.Account + updates map[int64]map[string]any +} + +func (r *grokQuotaHandlerAccountRepo) GetByID(_ context.Context, id int64) (*service.Account, error) { + if r.account != nil && r.account.ID == id { + return r.account, nil + } + return nil, service.ErrAccountNotFound +} + +func (r *grokQuotaHandlerAccountRepo) UpdateExtra(_ context.Context, id int64, updates map[string]any) error { + if r.updates == nil { + r.updates = make(map[int64]map[string]any) + } + r.updates[id] = updates + return nil +} + +type grokQuotaHandlerUpstream struct { + resp *http.Response + lastReq *http.Request + lastBody []byte +} + +func (u *grokQuotaHandlerUpstream) Do(req *http.Request, _ string, _ int64, _ int) (*http.Response, error) { + u.lastReq = req + if req.Body != nil { + u.lastBody, _ = io.ReadAll(req.Body) + } + return u.resp, nil +} + +func (u *grokQuotaHandlerUpstream) DoWithTLS( + req *http.Request, + proxyURL string, + accountID int64, + accountConcurrency int, + _ *tlsfingerprint.Profile, +) (*http.Response, error) { + return u.Do(req, proxyURL, accountID, accountConcurrency) +} + +func TestGrokOAuthHandlerQueryQuotaProbesUpstream(t *testing.T) { + gin.SetMode(gin.TestMode) + + repo := &grokQuotaHandlerAccountRepo{account: &service.Account{ + ID: 42, + Platform: service.PlatformGrok, + Type: service.AccountTypeOAuth, + Concurrency: 1, + Credentials: map[string]any{ + "access_token": "access-token", + "expires_at": time.Now().Add(time.Hour).UTC().Format(time.RFC3339), + }, + }} + upstream := &grokQuotaHandlerUpstream{resp: &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{ + "X-Ratelimit-Limit-Requests": []string{"10"}, + "X-Ratelimit-Remaining-Requests": []string{"8"}, + }, + Body: io.NopCloser(strings.NewReader(`{"id":"resp_probe"}`)), + }} + quotaService := service.NewGrokQuotaService(repo, nil, service.NewGrokTokenProvider(repo, nil), upstream) + handler := NewGrokOAuthHandler(nil, nil, quotaService) + + router := gin.New() + router.GET("/api/v1/admin/grok/accounts/:id/quota", handler.QueryQuota) + rec := httptest.NewRecorder() + req := httptest.NewRequest(http.MethodGet, "/api/v1/admin/grok/accounts/42/quota", nil) + router.ServeHTTP(rec, req) + + require.Equal(t, http.StatusOK, rec.Code) + require.Contains(t, rec.Body.String(), `"source":"active_probe"`) + require.Contains(t, rec.Body.String(), `"headers_observed":true`) + require.NotContains(t, rec.Body.String(), "access-token") + require.Equal(t, xai.DefaultBaseURL+"/responses", upstream.lastReq.URL.String()) + require.Equal(t, "Bearer access-token", upstream.lastReq.Header.Get("Authorization")) + require.Contains(t, string(upstream.lastBody), `"store":false`) + require.NotNil(t, repo.updates[42]) +} + +func TestGrokOAuthHandlerResetQuotaReturnsUnsupported(t *testing.T) { + gin.SetMode(gin.TestMode) + + repo := &grokQuotaHandlerAccountRepo{account: &service.Account{ + ID: 43, + Platform: service.PlatformGrok, + Type: service.AccountTypeOAuth, + }} + quotaService := service.NewGrokQuotaService(repo, nil, nil, nil) + handler := NewGrokOAuthHandler(nil, nil, quotaService) + + router := gin.New() + router.POST("/api/v1/admin/grok/accounts/:id/reset-quota", handler.ResetQuota) + rec := httptest.NewRecorder() + req := httptest.NewRequest(http.MethodPost, "/api/v1/admin/grok/accounts/43/reset-quota", nil) + router.ServeHTTP(rec, req) + + require.Equal(t, http.StatusNotImplemented, rec.Code) + require.Contains(t, rec.Body.String(), `"reason":"GROK_QUOTA_RESET_UNSUPPORTED"`) + require.NotContains(t, rec.Body.String(), "access-token") +} + +func TestGrokOAuthHandlerRuntimeSanityDoesNotExposeSecrets(t *testing.T) { + gin.SetMode(gin.TestMode) + t.Setenv(xai.EnvBaseURL, "http://127.0.0.1:8080/v1?access_token=secret") + t.Setenv(xai.EnvClientID, "client-secret-like-value") + + handler := NewGrokOAuthHandler(nil, nil, nil) + router := gin.New() + router.GET("/api/v1/admin/grok/runtime-sanity", handler.RuntimeSanity) + rec := httptest.NewRecorder() + req := httptest.NewRequest(http.MethodGet, "/api/v1/admin/grok/runtime-sanity", nil) + router.ServeHTTP(rec, req) + + require.Equal(t, http.StatusOK, rec.Code) + require.Contains(t, rec.Body.String(), `"public_gateway_scope":"responses_only"`) + require.Contains(t, rec.Body.String(), `"valid":false`) + require.NotContains(t, rec.Body.String(), "access_token") + require.NotContains(t, rec.Body.String(), "secret") + require.NotContains(t, rec.Body.String(), "client-secret-like-value") +} diff --git a/backend/internal/handler/admin/group_handler.go b/backend/internal/handler/admin/group_handler.go index 0468fc34e9..061611771d 100644 --- a/backend/internal/handler/admin/group_handler.go +++ b/backend/internal/handler/admin/group_handler.go @@ -84,7 +84,7 @@ func NewGroupHandler(adminService service.AdminService, dashboardService *servic type CreateGroupRequest struct { Name string `json:"name" binding:"required"` Description string `json:"description"` - Platform string `json:"platform" binding:"omitempty,oneof=anthropic openai gemini antigravity"` + Platform string `json:"platform" binding:"omitempty,oneof=anthropic openai gemini antigravity grok"` RateMultiplier float64 `json:"rate_multiplier"` IsExclusive bool `json:"is_exclusive"` SubscriptionType string `json:"subscription_type" binding:"omitempty,oneof=standard subscription"` @@ -124,7 +124,7 @@ type CreateGroupRequest struct { type UpdateGroupRequest struct { Name string `json:"name"` Description *string `json:"description"` - Platform string `json:"platform" binding:"omitempty,oneof=anthropic openai gemini antigravity"` + Platform string `json:"platform" binding:"omitempty,oneof=anthropic openai gemini antigravity grok"` RateMultiplier *float64 `json:"rate_multiplier"` IsExclusive *bool `json:"is_exclusive"` Status string `json:"status" binding:"omitempty,oneof=active inactive"` diff --git a/backend/internal/handler/admin/user_platform_quota_admin_test.go b/backend/internal/handler/admin/user_platform_quota_admin_test.go index fe33d36c5b..5689aedceb 100644 --- a/backend/internal/handler/admin/user_platform_quota_admin_test.go +++ b/backend/internal/handler/admin/user_platform_quota_admin_test.go @@ -112,9 +112,9 @@ func TestUpdateUserPlatformQuotas_Success(t *testing.T) { if repo.upsertCalls[0].userID != 42 || len(repo.upsertCalls[0].records) != 2 { t.Errorf("unexpected upsert call: %+v", repo.upsertCalls[0]) } - // 缓存失效:请求中 2 个 platform + 软删除的 2 个 platform(gemini, antigravity)= 4 次 - if len(cache.deleteCalls) != 4 { - t.Errorf("expected 4 cache delete calls, got %d: %+v", len(cache.deleteCalls), cache.deleteCalls) + // 缓存失效:请求中 2 个 platform + 软删除的 3 个 platform(gemini, antigravity, grok)= 5 次 + if len(cache.deleteCalls) != 5 { + t.Errorf("expected 5 cache delete calls, got %d: %+v", len(cache.deleteCalls), cache.deleteCalls) } } diff --git a/backend/internal/handler/endpoint.go b/backend/internal/handler/endpoint.go index 0d6f4b3cf6..e33d88241c 100644 --- a/backend/internal/handler/endpoint.go +++ b/backend/internal/handler/endpoint.go @@ -77,7 +77,7 @@ func DeriveUpstreamEndpoint(inbound, rawRequestPath, platform string) string { inbound = strings.TrimSpace(inbound) switch platform { - case service.PlatformOpenAI: + case service.PlatformOpenAI, service.PlatformGrok: if inbound == EndpointEmbeddings || inbound == EndpointImagesGenerations || inbound == EndpointImagesEdits { return inbound } diff --git a/backend/internal/handler/gateway_handler.go b/backend/internal/handler/gateway_handler.go index a46d688e9c..b65dedf9cc 100644 --- a/backend/internal/handler/gateway_handler.go +++ b/backend/internal/handler/gateway_handler.go @@ -25,6 +25,7 @@ import ( "github.com/Wei-Shaw/sub2api/internal/pkg/logger" "github.com/Wei-Shaw/sub2api/internal/pkg/openai" "github.com/Wei-Shaw/sub2api/internal/pkg/timezone" + "github.com/Wei-Shaw/sub2api/internal/pkg/xai" middleware2 "github.com/Wei-Shaw/sub2api/internal/server/middleware" "github.com/Wei-Shaw/sub2api/internal/service" @@ -1157,6 +1158,8 @@ func defaultModelIDsForPlatform(platform string) []string { ids = append(ids, model.ID) } return ids + case service.PlatformGrok: + return xai.DefaultModelIDs() default: ids := make([]string, 0, len(claude.DefaultModels)) for _, model := range claude.DefaultModels { diff --git a/backend/internal/handler/handler.go b/backend/internal/handler/handler.go index 95cac6274c..014cf7d2ba 100644 --- a/backend/internal/handler/handler.go +++ b/backend/internal/handler/handler.go @@ -17,6 +17,7 @@ type AdminHandlers struct { OpenAIOAuth *admin.OpenAIOAuthHandler GeminiOAuth *admin.GeminiOAuthHandler AntigravityOAuth *admin.AntigravityOAuthHandler + GrokOAuth *admin.GrokOAuthHandler Proxy *admin.ProxyHandler Redeem *admin.RedeemHandler Promo *admin.PromoHandler diff --git a/backend/internal/handler/openai_chat_completions.go b/backend/internal/handler/openai_chat_completions.go index c800397bb1..a91ebf8bd9 100644 --- a/backend/internal/handler/openai_chat_completions.go +++ b/backend/internal/handler/openai_chat_completions.go @@ -101,6 +101,7 @@ func (h *OpenAIGatewayHandler) ChatCompletions(c *gin.Context) { } subscription, _ := middleware2.GetSubscriptionFromContext(c) + requestPlatform := openAICompatibleRequestPlatform(apiKey) service.SetOpsLatencyMs(c, service.OpsAuthLatencyMsKey, time.Since(requestStart).Milliseconds()) routingStart := time.Now() @@ -144,6 +145,7 @@ func (h *OpenAIGatewayHandler) ChatCompletions(c *gin.Context) { service.OpenAIUpstreamTransportAny, service.OpenAIEndpointCapabilityChatCompletions, false, + requestPlatform, ) if err != nil { reqLog.Warn("openai_chat_completions.account_select_failed", diff --git a/backend/internal/handler/openai_gateway_handler.go b/backend/internal/handler/openai_gateway_handler.go index fbd13b66db..cc62121d94 100644 --- a/backend/internal/handler/openai_gateway_handler.go +++ b/backend/internal/handler/openai_gateway_handler.go @@ -97,6 +97,13 @@ func wrapUsageRecordTaskContext(parent context.Context, task service.UsageRecord } } +func openAICompatibleRequestPlatform(apiKey *service.APIKey) string { + if apiKey != nil && apiKey.Group != nil && apiKey.Group.Platform == service.PlatformGrok { + return service.PlatformGrok + } + return service.PlatformOpenAI +} + // NewOpenAIGatewayHandler creates a new OpenAIGatewayHandler func NewOpenAIGatewayHandler( gatewayService *service.OpenAIGatewayService, @@ -282,6 +289,7 @@ func (h *OpenAIGatewayHandler) Responses(c *gin.Context) { // Get subscription info (may be nil) subscription, _ := middleware2.GetSubscriptionFromContext(c) + requestPlatform := openAICompatibleRequestPlatform(apiKey) service.SetOpsLatencyMs(c, service.OpsAuthLatencyMsKey, time.Since(requestStart).Milliseconds()) routingStart := time.Now() @@ -332,6 +340,7 @@ func (h *OpenAIGatewayHandler) Responses(c *gin.Context) { service.OpenAIUpstreamTransportAny, service.OpenAIEndpointCapabilityChatCompletions, requireCompact, + requestPlatform, ) if err != nil { reqLog.Warn("openai.account_select_failed", @@ -708,6 +717,7 @@ func (h *OpenAIGatewayHandler) Messages(c *gin.Context) { } subscription, _ := middleware2.GetSubscriptionFromContext(c) + requestPlatform := openAICompatibleRequestPlatform(apiKey) service.SetOpsLatencyMs(c, service.OpsAuthLatencyMsKey, time.Since(requestStart).Milliseconds()) routingStart := time.Now() @@ -760,6 +770,7 @@ func (h *OpenAIGatewayHandler) Messages(c *gin.Context) { service.OpenAIUpstreamTransportAny, service.OpenAIEndpointCapabilityChatCompletions, false, + requestPlatform, ) if err != nil { reqLog.Warn("openai_messages.account_select_failed", @@ -1322,6 +1333,7 @@ func (h *OpenAIGatewayHandler) ResponsesWebSocket(c *gin.Context) { } subscription, _ := middleware2.GetSubscriptionFromContext(c) + requestPlatform := openAICompatibleRequestPlatform(apiKey) if err := h.billingCacheService.CheckBillingEligibility(ctx, apiKey.User, apiKey, apiKey.Group, subscription, service.QuotaPlatform(c.Request.Context(), apiKey)); err != nil { reqLog.Info("openai.websocket_billing_eligibility_check_failed", zap.Error(err)) closeOpenAIClientWS(wsConn, coderws.StatusPolicyViolation, "billing check failed") @@ -1350,6 +1362,7 @@ func (h *OpenAIGatewayHandler) ResponsesWebSocket(c *gin.Context) { service.OpenAIUpstreamTransportResponsesWebsocketV2, service.OpenAIEndpointCapabilityChatCompletions, false, + requestPlatform, ) if err != nil { reqLog.Warn("openai.websocket_account_select_failed", diff --git a/backend/internal/handler/openai_gateway_handler_test.go b/backend/internal/handler/openai_gateway_handler_test.go index fd3508837e..5886388a9e 100644 --- a/backend/internal/handler/openai_gateway_handler_test.go +++ b/backend/internal/handler/openai_gateway_handler_test.go @@ -1341,6 +1341,7 @@ func TestOpenAIResponsesWebSocket_FailoverOnUpstreamUsageLimitEvent(t *testing.T nil, nil, nil, + nil, ) cache := &concurrencyCacheMock{ @@ -1523,6 +1524,7 @@ func runOpenAIResponsesWebSocketUsageLogCase(t *testing.T, tc openAIResponsesWSU &service.DeferredService{}, nil, nil, + nil, channelSvc, nil, nil, diff --git a/backend/internal/handler/openai_images_failover_test.go b/backend/internal/handler/openai_images_failover_test.go index 40773eac43..e5101d0148 100644 --- a/backend/internal/handler/openai_images_failover_test.go +++ b/backend/internal/handler/openai_images_failover_test.go @@ -136,6 +136,7 @@ func TestOpenAIGatewayHandlerImages_ServerErrorFailsOverAndReturnsClearErrorWhen nil, nil, nil, + nil, ) billingService := service.NewBillingCacheService(nil, nil, nil, nil, nil, nil, cfg, nil) t.Cleanup(billingService.Stop) diff --git a/backend/internal/handler/wire.go b/backend/internal/handler/wire.go index 7d7f4dc622..090a734c9f 100644 --- a/backend/internal/handler/wire.go +++ b/backend/internal/handler/wire.go @@ -20,6 +20,7 @@ func ProvideAdminHandlers( openaiOAuthHandler *admin.OpenAIOAuthHandler, geminiOAuthHandler *admin.GeminiOAuthHandler, antigravityOAuthHandler *admin.AntigravityOAuthHandler, + grokOAuthHandler *admin.GrokOAuthHandler, proxyHandler *admin.ProxyHandler, redeemHandler *admin.RedeemHandler, promoHandler *admin.PromoHandler, @@ -53,6 +54,7 @@ func ProvideAdminHandlers( OpenAIOAuth: openaiOAuthHandler, GeminiOAuth: geminiOAuthHandler, AntigravityOAuth: antigravityOAuthHandler, + GrokOAuth: grokOAuthHandler, Proxy: proxyHandler, Redeem: redeemHandler, Promo: promoHandler, @@ -167,6 +169,7 @@ var ProviderSet = wire.NewSet( admin.NewOpenAIOAuthHandler, admin.NewGeminiOAuthHandler, admin.NewAntigravityOAuthHandler, + admin.NewGrokOAuthHandler, admin.NewProxyHandler, admin.NewRedeemHandler, admin.NewPromoHandler, diff --git a/backend/internal/model/error_passthrough_rule.go b/backend/internal/model/error_passthrough_rule.go index 620736cd87..aa202069bd 100644 --- a/backend/internal/model/error_passthrough_rule.go +++ b/backend/internal/model/error_passthrough_rule.go @@ -36,11 +36,12 @@ const ( PlatformOpenAI = "openai" PlatformGemini = "gemini" PlatformAntigravity = "antigravity" + PlatformGrok = "grok" ) // AllPlatforms 返回所有支持的平台列表 func AllPlatforms() []string { - return []string{PlatformAnthropic, PlatformOpenAI, PlatformGemini, PlatformAntigravity} + return []string{PlatformAnthropic, PlatformOpenAI, PlatformGemini, PlatformAntigravity, PlatformGrok} } // Validate 验证规则配置的有效性 diff --git a/backend/internal/pkg/xai/models.go b/backend/internal/pkg/xai/models.go new file mode 100644 index 0000000000..0d289274fb --- /dev/null +++ b/backend/internal/pkg/xai/models.go @@ -0,0 +1,46 @@ +package xai + +// Model describes an xAI model in OpenAI-compatible /models shape. +type Model struct { + ID string `json:"id"` + Object string `json:"object"` + Created int64 `json:"created,omitempty"` + OwnedBy string `json:"owned_by"` + DisplayName string `json:"display_name,omitempty"` +} + +var defaultModels = []Model{ + {ID: "grok-4.3", Object: "model", OwnedBy: "xai", DisplayName: "Grok 4.3"}, + {ID: "grok-build-0.1", Object: "model", OwnedBy: "xai", DisplayName: "Grok Build 0.1"}, + {ID: "grok-4.20-0309-reasoning", Object: "model", OwnedBy: "xai", DisplayName: "Grok 4.20 Reasoning"}, + {ID: "grok-4.20-0309-non-reasoning", Object: "model", OwnedBy: "xai", DisplayName: "Grok 4.20 Non Reasoning"}, + {ID: "grok-4.20-multi-agent-0309", Object: "model", OwnedBy: "xai", DisplayName: "Grok 4.20 Multi Agent"}, +} + +func DefaultModels() []Model { + out := make([]Model, len(defaultModels)) + copy(out, defaultModels) + return out +} + +func DefaultModelIDs() []string { + models := DefaultModels() + ids := make([]string, 0, len(models)) + for _, model := range models { + ids = append(ids, model.ID) + } + return ids +} + +func DefaultModelMapping() map[string]string { + mapping := make(map[string]string, len(defaultModels)+3) + for _, model := range defaultModels { + mapping[model.ID] = model.ID + } + mapping["grok"] = "grok-4.3" + mapping["grok-latest"] = "grok-4.3" + mapping["grok-build"] = "grok-build-0.1" + mapping["grok-4.20-reasoning"] = "grok-4.20-0309-reasoning" + mapping["grok-4.20-non-reasoning"] = "grok-4.20-0309-non-reasoning" + return mapping +} diff --git a/backend/internal/pkg/xai/oauth.go b/backend/internal/pkg/xai/oauth.go new file mode 100644 index 0000000000..449b5cb865 --- /dev/null +++ b/backend/internal/pkg/xai/oauth.go @@ -0,0 +1,448 @@ +package xai + +import ( + "crypto/rand" + "crypto/sha256" + "encoding/base64" + "encoding/hex" + "fmt" + "net/url" + "os" + "strings" + "sync" + "time" + + "github.com/Wei-Shaw/sub2api/internal/util/logredact" + "github.com/Wei-Shaw/sub2api/internal/util/urlvalidator" +) + +const ( + OAuthIssuer = "https://auth.x.ai" + DiscoveryURL = OAuthIssuer + "/.well-known/openid-configuration" + DefaultAuthorizeURL = OAuthIssuer + "/oauth2/authorize" + DefaultTokenURL = OAuthIssuer + "/oauth2/token" + DefaultBaseURL = "https://api.x.ai/v1" + DefaultCLIBaseURL = "https://cli-chat-proxy.grok.com/v1" + DefaultClientID = "b1a00492-073a-47ea-816f-4c329264a828" + DefaultScope = "openid profile email offline_access grok-cli:access api:access" + DefaultRedirectURI = "http://127.0.0.1:56121/callback" + SessionTTL = 30 * time.Minute + + EnvAuthorizeURL = "XAI_OAUTH_AUTHORIZE_URL" + EnvTokenURL = "XAI_OAUTH_TOKEN_URL" + EnvClientID = "XAI_OAUTH_CLIENT_ID" + EnvScope = "XAI_OAUTH_SCOPE" + EnvRedirectURI = "XAI_OAUTH_REDIRECT_URI" + EnvBaseURL = "XAI_BASE_URL" + EnvAllowUnsafeURLOverrides = "XAI_ALLOW_UNSAFE_URL_OVERRIDES" + EnvUnsafeAllowHighConcurrency = "XAI_GROK_UNSAFE_ALLOW_CONCURRENCY_GT_ONE" +) + +var ( + oauthEndpointAllowedHosts = []string{"x.ai", "*.x.ai"} + baseURLAllowedHosts = []string{"api.x.ai", "cli-chat-proxy.grok.com"} +) + +// OAuthSession stores one PKCE OAuth flow. +type OAuthSession struct { + State string `json:"state"` + CodeVerifier string `json:"code_verifier"` + CodeChallenge string `json:"code_challenge"` + ClientID string `json:"client_id,omitempty"` + Scope string `json:"scope,omitempty"` + ProxyURL string `json:"proxy_url,omitempty"` + RedirectURI string `json:"redirect_uri"` + CreatedAt time.Time `json:"created_at"` +} + +// SessionStore manages xAI OAuth sessions in memory. +type SessionStore struct { + mu sync.RWMutex + sessions map[string]*OAuthSession + stopOnce sync.Once + stopCh chan struct{} +} + +func NewSessionStore() *SessionStore { + store := &SessionStore{ + sessions: make(map[string]*OAuthSession), + stopCh: make(chan struct{}), + } + go store.cleanup() + return store +} + +func (s *SessionStore) Set(sessionID string, session *OAuthSession) { + s.mu.Lock() + defer s.mu.Unlock() + s.sessions[sessionID] = session +} + +func (s *SessionStore) Get(sessionID string) (*OAuthSession, bool) { + s.mu.RLock() + defer s.mu.RUnlock() + session, ok := s.sessions[sessionID] + if !ok { + return nil, false + } + if time.Since(session.CreatedAt) > SessionTTL { + return nil, false + } + return session, true +} + +func (s *SessionStore) Delete(sessionID string) { + s.mu.Lock() + defer s.mu.Unlock() + delete(s.sessions, sessionID) +} + +func (s *SessionStore) Stop() { + s.stopOnce.Do(func() { + close(s.stopCh) + }) +} + +func (s *SessionStore) cleanup() { + ticker := time.NewTicker(5 * time.Minute) + defer ticker.Stop() + for { + select { + case <-s.stopCh: + return + case <-ticker.C: + s.mu.Lock() + for id, session := range s.sessions { + if time.Since(session.CreatedAt) > SessionTTL { + delete(s.sessions, id) + } + } + s.mu.Unlock() + } + } +} + +func EffectiveAuthorizeURL() string { + return envOrDefault(EnvAuthorizeURL, DefaultAuthorizeURL) +} + +func ValidatedAuthorizeURL() (string, error) { + return ValidateOAuthEndpointURL(EffectiveAuthorizeURL()) +} + +func EffectiveTokenURL() string { + return envOrDefault(EnvTokenURL, DefaultTokenURL) +} + +func ValidatedTokenURL() (string, error) { + return ValidateOAuthEndpointURL(EffectiveTokenURL()) +} + +func EffectiveClientID() string { + return envOrDefault(EnvClientID, DefaultClientID) +} + +func EffectiveScope() string { + return envOrDefault(EnvScope, DefaultScope) +} + +func EffectiveRedirectURI(override string) string { + if trimmed := strings.TrimSpace(override); trimmed != "" { + return trimmed + } + return envOrDefault(EnvRedirectURI, DefaultRedirectURI) +} + +func EffectiveBaseURL(override string) string { + if trimmed := strings.TrimSpace(override); trimmed != "" { + return strings.TrimRight(trimmed, "/") + } + return strings.TrimRight(envOrDefault(EnvBaseURL, DefaultBaseURL), "/") +} + +func ValidatedBaseURL(override string) (string, error) { + return ValidateBaseURL(EffectiveBaseURL(override)) +} + +type RuntimeSanityCheck struct { + Value string `json:"value"` + Valid bool `json:"valid"` + Error string `json:"error,omitempty"` + IsDefault bool `json:"is_default,omitempty"` +} + +type RuntimeSanityReport struct { + BaseURL RuntimeSanityCheck `json:"base_url"` + OAuthAuthorizeURL RuntimeSanityCheck `json:"oauth_authorize_url"` + OAuthTokenURL RuntimeSanityCheck `json:"oauth_token_url"` + OAuthRedirectURI RuntimeSanityCheck `json:"oauth_redirect_uri"` + UnsafeURLOverrides bool `json:"unsafe_url_overrides"` + UnsafeHighConcurrency bool `json:"unsafe_high_concurrency"` + PublicGatewayScope string `json:"public_gateway_scope"` + ProxyPolicy string `json:"proxy_policy"` +} + +func RuntimeSanity() RuntimeSanityReport { + return RuntimeSanityReport{ + BaseURL: runtimeSanityCheck(EffectiveBaseURL(""), EnvBaseURL, ValidatedBaseURL), + OAuthAuthorizeURL: runtimeSanityCheck(EffectiveAuthorizeURL(), EnvAuthorizeURL, func(string) (string, error) { return ValidatedAuthorizeURL() }), + OAuthTokenURL: runtimeSanityCheck(EffectiveTokenURL(), EnvTokenURL, func(string) (string, error) { return ValidatedTokenURL() }), + OAuthRedirectURI: runtimeSanityCheck(EffectiveRedirectURI(""), EnvRedirectURI, validateRedirectURI), + UnsafeURLOverrides: AllowUnsafeURLOverrides(), + UnsafeHighConcurrency: AllowUnsafeHighConcurrency(), + PublicGatewayScope: "responses_only", + ProxyPolicy: "account_proxy_optional; upstream URL allowlists enforced unless unsafe overrides are enabled", + } +} + +func runtimeSanityCheck(value string, envKey string, validate func(string) (string, error)) RuntimeSanityCheck { + normalized, err := validate(value) + check := RuntimeSanityCheck{ + Value: sanitizeRuntimeURLValue(normalized), + Valid: err == nil, + IsDefault: strings.TrimSpace(os.Getenv(envKey)) == "", + } + if err != nil { + check.Value = sanitizeRuntimeURLValue(value) + check.Error = sanitizeRuntimeError(err.Error(), value) + } + return check +} + +func validateRedirectURI(raw string) (string, error) { + return urlvalidator.ValidateURLFormat(raw, true) +} + +func sanitizeRuntimeURLValue(raw string) string { + trimmed := strings.TrimSpace(raw) + if trimmed == "" { + return "" + } + parsed, err := url.Parse(trimmed) + if err != nil || parsed.Scheme == "" || parsed.Host == "" { + return trimmed + } + parsed.User = nil + parsed.RawQuery = "" + parsed.Fragment = "" + return strings.TrimRight(parsed.String(), "/") +} + +func sanitizeRuntimeError(rawErr string, rawValue string) string { + redacted := logredact.RedactText(rawErr) + trimmedValue := strings.TrimSpace(rawValue) + if trimmedValue == "" { + return redacted + } + sanitizedValue := sanitizeRuntimeURLValue(trimmedValue) + redacted = strings.ReplaceAll(redacted, trimmedValue, sanitizedValue) + redacted = strings.ReplaceAll(redacted, logredact.RedactText(trimmedValue), sanitizedValue) + return redacted +} + +func ValidateOAuthEndpointURL(raw string) (string, error) { + if AllowUnsafeURLOverrides() { + return urlvalidator.ValidateURLFormat(raw, true) + } + return urlvalidator.ValidateHTTPSURL(raw, urlvalidator.ValidationOptions{ + AllowedHosts: oauthEndpointAllowedHosts, + RequireAllowlist: true, + AllowPrivate: false, + }) +} + +func ValidateBaseURL(raw string) (string, error) { + if AllowUnsafeURLOverrides() { + return urlvalidator.ValidateURLFormat(raw, true) + } + normalized, err := urlvalidator.ValidateHTTPSURL(raw, urlvalidator.ValidationOptions{ + AllowedHosts: baseURLAllowedHosts, + RequireAllowlist: true, + AllowPrivate: false, + }) + if err != nil { + return "", err + } + return normalizeKnownBaseURLPath(normalized) +} + +func normalizeKnownBaseURLPath(raw string) (string, error) { + parsed, err := url.Parse(raw) + if err != nil || parsed.Scheme == "" || parsed.Host == "" { + return "", fmt.Errorf("invalid url: %s", raw) + } + path := strings.TrimRight(parsed.Path, "/") + if path == "" { + parsed.Path = "/v1" + parsed.RawPath = "" + return strings.TrimRight(parsed.String(), "/"), nil + } + if path != "/v1" { + return "", fmt.Errorf("base URL path must be /v1") + } + parsed.Path = path + parsed.RawPath = "" + return strings.TrimRight(parsed.String(), "/"), nil +} + +func AllowUnsafeURLOverrides() bool { + return envBool(EnvAllowUnsafeURLOverrides) +} + +func AllowUnsafeHighConcurrency() bool { + return envBool(EnvUnsafeAllowHighConcurrency) +} + +func envOrDefault(key, fallback string) string { + if value := strings.TrimSpace(os.Getenv(key)); value != "" { + return value + } + return fallback +} + +func envBool(key string) bool { + switch strings.ToLower(strings.TrimSpace(os.Getenv(key))) { + case "1", "true", "yes", "y", "on": + return true + default: + return false + } +} + +func GenerateRandomBytes(n int) ([]byte, error) { + b := make([]byte, n) + if _, err := rand.Read(b); err != nil { + return nil, err + } + return b, nil +} + +func GenerateState() (string, error) { + bytes, err := GenerateRandomBytes(32) + if err != nil { + return "", err + } + return hex.EncodeToString(bytes), nil +} + +func GenerateNonce() (string, error) { + bytes, err := GenerateRandomBytes(16) + if err != nil { + return "", err + } + return hex.EncodeToString(bytes), nil +} + +func GenerateSessionID() (string, error) { + bytes, err := GenerateRandomBytes(16) + if err != nil { + return "", err + } + return hex.EncodeToString(bytes), nil +} + +func GenerateCodeVerifier() (string, error) { + bytes, err := GenerateRandomBytes(32) + if err != nil { + return "", err + } + return base64URLEncode(bytes), nil +} + +func GenerateCodeChallenge(verifier string) string { + hash := sha256.Sum256([]byte(verifier)) + return base64URLEncode(hash[:]) +} + +func base64URLEncode(data []byte) string { + return strings.TrimRight(base64.URLEncoding.EncodeToString(data), "=") +} + +func BuildAuthorizationURL(state, codeChallenge, redirectURI, nonce string) (string, error) { + redirectURI = EffectiveRedirectURI(redirectURI) + authorizeURL, err := ValidatedAuthorizeURL() + if err != nil { + return "", fmt.Errorf("invalid authorize url: %w", err) + } + + params := url.Values{} + params.Set("response_type", "code") + params.Set("client_id", EffectiveClientID()) + params.Set("redirect_uri", redirectURI) + params.Set("scope", EffectiveScope()) + params.Set("state", state) + params.Set("nonce", nonce) + params.Set("code_challenge", codeChallenge) + params.Set("code_challenge_method", "S256") + params.Set("plan", "generic") + params.Set("referrer", "sub2api") + + return fmt.Sprintf("%s?%s", authorizeURL, params.Encode()), nil +} + +// AuthorizationInput is a parsed manual OAuth callback input. +type AuthorizationInput struct { + Code string + State string + RequiresState bool +} + +// ParseAuthorizationInput accepts a full callback URL, query string, or bare code. +func ParseAuthorizationInput(raw string) AuthorizationInput { + trimmed := strings.TrimSpace(raw) + if trimmed == "" { + return AuthorizationInput{} + } + + if parsed, err := url.Parse(trimmed); err == nil && parsed != nil { + values := parsed.Query() + if code := strings.TrimSpace(values.Get("code")); code != "" { + return AuthorizationInput{ + Code: code, + State: strings.TrimSpace(values.Get("state")), + RequiresState: true, + } + } + } + + queryCandidate := strings.TrimPrefix(trimmed, "?") + if strings.Contains(queryCandidate, "=") { + 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")), + RequiresState: true, + } + } + } + } + + return AuthorizationInput{Code: trimmed} +} + +func BuildResponsesURL(baseURL string) (string, error) { + validatedBaseURL, err := ValidatedBaseURL(baseURL) + if err != nil { + return "", fmt.Errorf("invalid base url: %w", err) + } + return validatedBaseURL + "/responses", nil +} + +func BuildChatCompletionsURL(baseURL string) (string, error) { + validatedBaseURL, err := ValidatedBaseURL(baseURL) + if err != nil { + return "", fmt.Errorf("invalid base url: %w", err) + } + return validatedBaseURL + "/chat/completions", nil +} + +// TokenResponse represents xAI OAuth token responses. +type TokenResponse struct { + AccessToken string `json:"access_token"` + RefreshToken string `json:"refresh_token,omitempty"` + IDToken string `json:"id_token,omitempty"` + TokenType string `json:"token_type,omitempty"` + ExpiresIn int64 `json:"expires_in,omitempty"` + Scope string `json:"scope,omitempty"` +} diff --git a/backend/internal/pkg/xai/oauth_test.go b/backend/internal/pkg/xai/oauth_test.go new file mode 100644 index 0000000000..fc48182fe7 --- /dev/null +++ b/backend/internal/pkg/xai/oauth_test.go @@ -0,0 +1,195 @@ +//go:build unit + +package xai + +import ( + "net/url" + "testing" + + "github.com/stretchr/testify/require" +) + +func TestParseAuthorizationInput(t *testing.T) { + t.Parallel() + + tests := []struct { + 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", + wantRequiresState: true, + }, + { + 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", + raw: "abc123", + wantCode: "abc123", + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + t.Parallel() + 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) + }) + } +} + +func TestBuildAuthorizationURLIncludesHermesCompatibleParameters(t *testing.T) { + t.Setenv(EnvAuthorizeURL, "https://auth.example.test/oauth2/authorize") + t.Setenv(EnvClientID, "client-id") + t.Setenv(EnvScope, "openid profile offline_access api:access") + t.Setenv(EnvAllowUnsafeURLOverrides, "true") + + authURL, err := BuildAuthorizationURL("state", "challenge", "http://127.0.0.1:56121/callback", "nonce") + require.NoError(t, err) + parsed, err := url.Parse(authURL) + require.NoError(t, err) + + values := parsed.Query() + require.Equal(t, "https", parsed.Scheme) + require.Equal(t, "auth.example.test", parsed.Host) + require.Equal(t, "/oauth2/authorize", parsed.Path) + require.Equal(t, "code", values.Get("response_type")) + require.Equal(t, "client-id", values.Get("client_id")) + require.Equal(t, "http://127.0.0.1:56121/callback", values.Get("redirect_uri")) + require.Equal(t, "openid profile offline_access api:access", values.Get("scope")) + require.Equal(t, "state", values.Get("state")) + require.Equal(t, "nonce", values.Get("nonce")) + require.Equal(t, "challenge", values.Get("code_challenge")) + require.Equal(t, "S256", values.Get("code_challenge_method")) + require.Equal(t, "generic", values.Get("plan")) + require.Equal(t, "sub2api", values.Get("referrer")) +} + +func TestValidateXAIURLsAllowOfficialOAuthAndGatewayHosts(t *testing.T) { + authorizeURL, err := ValidateOAuthEndpointURL(DefaultAuthorizeURL) + require.NoError(t, err) + require.Equal(t, DefaultAuthorizeURL, authorizeURL) + + tokenURL, err := ValidateOAuthEndpointURL(DefaultTokenURL) + require.NoError(t, err) + require.Equal(t, DefaultTokenURL, tokenURL) + + baseURL, err := ValidateBaseURL(DefaultBaseURL) + require.NoError(t, err) + require.Equal(t, DefaultBaseURL, baseURL) + + cliBaseURL, err := ValidateBaseURL(DefaultCLIBaseURL) + require.NoError(t, err) + require.Equal(t, DefaultCLIBaseURL, cliBaseURL) + + baseURLNoPath, err := ValidateBaseURL("https://api.x.ai") + require.NoError(t, err) + require.Equal(t, DefaultBaseURL, baseURLNoPath) + + chatURL, err := BuildChatCompletionsURL(DefaultCLIBaseURL + "/") + require.NoError(t, err) + require.Equal(t, DefaultCLIBaseURL+"/chat/completions", chatURL) +} + +func TestValidateXAIURLsRejectArbitraryHostsByDefault(t *testing.T) { + _, err := ValidateOAuthEndpointURL("https://auth.example.test/oauth2/token") + require.Error(t, err) + + _, err = ValidateBaseURL("https://xai.test/v1") + require.Error(t, err) + + _, err = ValidateBaseURL("http://127.0.0.1:8080/v1") + require.Error(t, err) + + _, err = ValidateBaseURL("https://api.x.ai/custom") + require.Error(t, err) +} + +func TestValidateXAIURLsAllowUnsafeDevOverride(t *testing.T) { + t.Setenv(EnvAllowUnsafeURLOverrides, "true") + + tokenURL, err := ValidateOAuthEndpointURL("http://127.0.0.1:8080/oauth2/token") + require.NoError(t, err) + require.Equal(t, "http://127.0.0.1:8080/oauth2/token", tokenURL) + + baseURL, err := ValidateBaseURL("http://127.0.0.1:8080/v1/") + require.NoError(t, err) + require.Equal(t, "http://127.0.0.1:8080/v1", baseURL) +} + +func TestRuntimeSanityReportsSafeDefaults(t *testing.T) { + t.Setenv(EnvBaseURL, "") + t.Setenv(EnvAuthorizeURL, "") + t.Setenv(EnvTokenURL, "") + t.Setenv(EnvRedirectURI, "") + t.Setenv(EnvAllowUnsafeURLOverrides, "") + t.Setenv(EnvUnsafeAllowHighConcurrency, "") + + report := RuntimeSanity() + require.True(t, report.BaseURL.Valid) + require.Equal(t, DefaultBaseURL, report.BaseURL.Value) + require.True(t, report.BaseURL.IsDefault) + require.True(t, report.OAuthAuthorizeURL.Valid) + require.True(t, report.OAuthTokenURL.Valid) + require.True(t, report.OAuthRedirectURI.Valid) + require.False(t, report.UnsafeURLOverrides) + require.False(t, report.UnsafeHighConcurrency) + require.Equal(t, "responses_only", report.PublicGatewayScope) + require.Contains(t, report.ProxyPolicy, "account_proxy_optional") +} + +func TestRuntimeSanityReportsInvalidOverridesWithoutSecrets(t *testing.T) { + t.Setenv(EnvBaseURL, "http://127.0.0.1:8080/v1?access_token=secret") + t.Setenv(EnvAuthorizeURL, "https://auth.example.test/oauth2/authorize") + t.Setenv(EnvTokenURL, "https://auth.example.test/oauth2/token") + t.Setenv(EnvRedirectURI, "not a url") + t.Setenv(EnvClientID, "client-secret-like-value") + t.Setenv(EnvAllowUnsafeURLOverrides, "") + + report := RuntimeSanity() + require.False(t, report.BaseURL.Valid) + require.False(t, report.BaseURL.IsDefault) + require.Contains(t, report.BaseURL.Error, "invalid url") + require.NotContains(t, report.BaseURL.Value, "secret") + require.False(t, report.OAuthAuthorizeURL.Valid) + require.False(t, report.OAuthTokenURL.Valid) + require.False(t, report.OAuthRedirectURI.Valid) + require.NotContains(t, report.ProxyPolicy, "client-secret-like-value") +} + +func TestDefaultModelMappingIncludesGrokAliases(t *testing.T) { + t.Parallel() + + mapping := DefaultModelMapping() + require.Equal(t, "grok-4.3", mapping["grok"]) + require.Equal(t, "grok-4.3", mapping["grok-latest"]) + require.Equal(t, "grok-build-0.1", mapping["grok-build"]) + require.Equal(t, "grok-4.20-0309-reasoning", mapping["grok-4.20-reasoning"]) + require.Equal(t, "grok-4.20-0309-non-reasoning", mapping["grok-4.20-non-reasoning"]) + require.Equal(t, "grok-4.20-multi-agent-0309", mapping["grok-4.20-multi-agent-0309"]) +} diff --git a/backend/internal/pkg/xai/quota.go b/backend/internal/pkg/xai/quota.go new file mode 100644 index 0000000000..1387c5bdad --- /dev/null +++ b/backend/internal/pkg/xai/quota.go @@ -0,0 +1,181 @@ +package xai + +import ( + "net/http" + "strconv" + "strings" + "time" +) + +type QuotaWindow struct { + Limit *int64 `json:"limit,omitempty"` + Remaining *int64 `json:"remaining,omitempty"` + ResetUnix *int64 `json:"reset_unix,omitempty"` + ResetAt string `json:"reset_at,omitempty"` +} + +type QuotaSnapshot struct { + Requests *QuotaWindow `json:"requests,omitempty"` + Tokens *QuotaWindow `json:"tokens,omitempty"` + RetryAfterSeconds *int `json:"retry_after_seconds,omitempty"` + SubscriptionTier string `json:"subscription_tier,omitempty"` + EntitlementStatus string `json:"entitlement_status,omitempty"` + StatusCode int `json:"status_code,omitempty"` + Headers map[string]string `json:"headers,omitempty"` + HeadersObserved bool `json:"headers_observed"` + ObservationSource string `json:"observation_source,omitempty"` + LastProbeAt string `json:"last_probe_at,omitempty"` + LastHeadersSeenAt string `json:"last_headers_seen_at,omitempty"` + UpdatedAt string `json:"updated_at"` +} + +func (s *QuotaSnapshot) HasObservedHeaders() bool { + if s == nil { + return false + } + return s.HeadersObserved || + s.Requests != nil || + s.Tokens != nil || + s.RetryAfterSeconds != nil || + s.SubscriptionTier != "" || + s.EntitlementStatus != "" || + len(s.Headers) > 0 +} + +var quotaHeaderAllowlist = []string{ + "x-ratelimit-limit-requests", + "x-ratelimit-remaining-requests", + "x-ratelimit-reset-requests", + "x-ratelimit-limit-tokens", + "x-ratelimit-remaining-tokens", + "x-ratelimit-reset-tokens", + "retry-after", + "x-subscription-tier", + "xai-subscription-tier", + "x-entitlement-status", + "xai-entitlement-status", +} + +func ParseQuotaHeaders(headers http.Header, statusCode int) *QuotaSnapshot { + return parseQuotaHeaders(headers, statusCode, "", false) +} + +func ObserveQuotaHeaders(headers http.Header, statusCode int, source string) *QuotaSnapshot { + return parseQuotaHeaders(headers, statusCode, source, true) +} + +func parseQuotaHeaders(headers http.Header, statusCode int, source string, keepEmpty bool) *QuotaSnapshot { + if headers == nil && !keepEmpty { + return nil + } + now := time.Now().UTC().Format(time.RFC3339) + snapshot := &QuotaSnapshot{ + Requests: parseQuotaWindow(headers, "requests"), + Tokens: parseQuotaWindow(headers, "tokens"), + StatusCode: statusCode, + Headers: make(map[string]string), + ObservationSource: strings.TrimSpace(source), + UpdatedAt: now, + } + if snapshot.ObservationSource == "active_probe" { + snapshot.LastProbeAt = now + } + if retryAfter := parseRetryAfter(headers.Get("retry-after")); retryAfter != nil { + snapshot.RetryAfterSeconds = retryAfter + } + snapshot.SubscriptionTier = firstHeader(headers, "xai-subscription-tier", "x-subscription-tier") + snapshot.EntitlementStatus = firstHeader(headers, "xai-entitlement-status", "x-entitlement-status") + + for _, name := range quotaHeaderAllowlist { + if value := strings.TrimSpace(headers.Get(name)); value != "" { + snapshot.Headers[name] = value + } + } + + if snapshot.Requests == nil && + snapshot.Tokens == nil && + snapshot.RetryAfterSeconds == nil && + snapshot.SubscriptionTier == "" && + snapshot.EntitlementStatus == "" && + len(snapshot.Headers) == 0 { + if keepEmpty { + return snapshot + } + return nil + } + snapshot.HeadersObserved = true + snapshot.LastHeadersSeenAt = now + return snapshot +} + +func parseQuotaWindow(headers http.Header, dimension string) *QuotaWindow { + window := &QuotaWindow{ + Limit: parseInt64Ptr(headers.Get("x-ratelimit-limit-" + dimension)), + Remaining: parseInt64Ptr(headers.Get("x-ratelimit-remaining-" + dimension)), + } + if reset := parseResetHeader(headers.Get("x-ratelimit-reset-" + dimension)); reset != nil { + window.ResetUnix = reset + window.ResetAt = time.Unix(*reset, 0).UTC().Format(time.RFC3339) + } + if window.Limit == nil && window.Remaining == nil && window.ResetUnix == nil { + return nil + } + return window +} + +func parseResetHeader(raw string) *int64 { + raw = strings.TrimSpace(raw) + if raw == "" { + return nil + } + if value, err := strconv.ParseInt(raw, 10, 64); err == nil { + if value > 1_000_000_000_000 { + value = value / 1000 + } + return &value + } + if t, err := time.Parse(time.RFC3339, raw); err == nil { + value := t.Unix() + return &value + } + return nil +} + +func parseRetryAfter(raw string) *int { + raw = strings.TrimSpace(raw) + if raw == "" { + return nil + } + if value, err := strconv.Atoi(raw); err == nil { + return &value + } + if t, err := http.ParseTime(raw); err == nil { + seconds := int(time.Until(t).Seconds()) + if seconds < 0 { + seconds = 0 + } + return &seconds + } + return nil +} + +func parseInt64Ptr(raw string) *int64 { + raw = strings.TrimSpace(raw) + if raw == "" { + return nil + } + value, err := strconv.ParseInt(raw, 10, 64) + if err != nil { + return nil + } + return &value +} + +func firstHeader(headers http.Header, names ...string) string { + for _, name := range names { + if value := strings.TrimSpace(headers.Get(name)); value != "" { + return value + } + } + return "" +} diff --git a/backend/internal/pkg/xai/quota_test.go b/backend/internal/pkg/xai/quota_test.go new file mode 100644 index 0000000000..593edc8eb9 --- /dev/null +++ b/backend/internal/pkg/xai/quota_test.go @@ -0,0 +1,63 @@ +//go:build unit + +package xai + +import ( + "net/http" + "testing" + + "github.com/stretchr/testify/require" +) + +func TestParseQuotaHeaders(t *testing.T) { + t.Parallel() + + headers := http.Header{} + headers.Set("x-ratelimit-limit-requests", "100") + headers.Set("x-ratelimit-remaining-requests", "25") + headers.Set("x-ratelimit-reset-requests", "1893456000") + headers.Set("x-ratelimit-limit-tokens", "1000000") + headers.Set("x-ratelimit-remaining-tokens", "750000") + headers.Set("retry-after", "60") + headers.Set("xai-subscription-tier", "supergrok") + headers.Set("xai-entitlement-status", "active") + headers.Set("authorization", "should-not-be-copied") + + snapshot := ParseQuotaHeaders(headers, http.StatusTooManyRequests) + require.NotNil(t, snapshot) + require.Equal(t, http.StatusTooManyRequests, snapshot.StatusCode) + require.True(t, snapshot.HeadersObserved) + require.NotEmpty(t, snapshot.LastHeadersSeenAt) + require.Equal(t, int64(100), *snapshot.Requests.Limit) + require.Equal(t, int64(25), *snapshot.Requests.Remaining) + require.Equal(t, int64(1893456000), *snapshot.Requests.ResetUnix) + require.Equal(t, "2030-01-01T00:00:00Z", snapshot.Requests.ResetAt) + require.Equal(t, int64(1000000), *snapshot.Tokens.Limit) + require.Equal(t, int64(750000), *snapshot.Tokens.Remaining) + require.Equal(t, 60, *snapshot.RetryAfterSeconds) + require.Equal(t, "supergrok", snapshot.SubscriptionTier) + require.Equal(t, "active", snapshot.EntitlementStatus) + require.Contains(t, snapshot.Headers, "x-ratelimit-limit-requests") + require.NotContains(t, snapshot.Headers, "authorization") +} + +func TestParseQuotaHeadersReturnsNilForMissingHeaders(t *testing.T) { + t.Parallel() + + require.Nil(t, ParseQuotaHeaders(http.Header{}, http.StatusOK)) +} + +func TestObserveQuotaHeadersRecordsNoHeaderProbe(t *testing.T) { + t.Parallel() + + snapshot := ObserveQuotaHeaders(http.Header{}, http.StatusOK, "active_probe") + require.NotNil(t, snapshot) + require.False(t, snapshot.HeadersObserved) + require.Equal(t, http.StatusOK, snapshot.StatusCode) + require.Equal(t, "active_probe", snapshot.ObservationSource) + require.NotEmpty(t, snapshot.LastProbeAt) + require.Empty(t, snapshot.LastHeadersSeenAt) + require.Empty(t, snapshot.Headers) + require.Nil(t, snapshot.Requests) + require.Nil(t, snapshot.Tokens) +} diff --git a/backend/internal/repository/grok_oauth_client.go b/backend/internal/repository/grok_oauth_client.go new file mode 100644 index 0000000000..435ced5a65 --- /dev/null +++ b/backend/internal/repository/grok_oauth_client.go @@ -0,0 +1,112 @@ +package repository + +import ( + "context" + "net/http" + "net/url" + "strings" + "time" + + infraerrors "github.com/Wei-Shaw/sub2api/internal/pkg/errors" + "github.com/Wei-Shaw/sub2api/internal/pkg/xai" + "github.com/Wei-Shaw/sub2api/internal/service" + "github.com/Wei-Shaw/sub2api/internal/util/logredact" + "github.com/imroc/req/v3" +) + +type grokOAuthClient struct { + tokenURL string +} + +func NewGrokOAuthClient() service.GrokOAuthClient { + return &grokOAuthClient{tokenURL: xai.EffectiveTokenURL()} +} + +func (c *grokOAuthClient) ExchangeCode(ctx context.Context, code, codeVerifier, redirectURI, proxyURL, clientID string) (*xai.TokenResponse, error) { + client, err := createGrokReqClient(proxyURL) + if err != nil { + return nil, infraerrors.Newf(http.StatusBadGateway, "GROK_OAUTH_CLIENT_INIT_FAILED", "create HTTP client: %v", err) + } + + clientID = strings.TrimSpace(clientID) + if clientID == "" { + clientID = xai.EffectiveClientID() + } + + formData := url.Values{} + formData.Set("grant_type", "authorization_code") + formData.Set("client_id", clientID) + formData.Set("code", code) + formData.Set("redirect_uri", xai.EffectiveRedirectURI(redirectURI)) + formData.Set("code_verifier", codeVerifier) + + var tokenResp xai.TokenResponse + resp, err := client.R(). + SetContext(ctx). + SetHeader("User-Agent", "sub2api-grok-oauth/1.0"). + SetFormDataFromValues(formData). + SetSuccessResult(&tokenResp). + Post(c.tokenURL) + if err != nil { + return nil, infraerrors.Newf(http.StatusBadGateway, "GROK_OAUTH_REQUEST_FAILED", "request failed: %v", err) + } + if !resp.IsSuccessState() { + return nil, grokOAuthStatusError("GROK_OAUTH_TOKEN_EXCHANGE_FAILED", "token exchange failed", resp) + } + return &tokenResp, nil +} + +func (c *grokOAuthClient) RefreshToken(ctx context.Context, refreshToken, proxyURL, clientID string) (*xai.TokenResponse, error) { + client, err := createGrokReqClient(proxyURL) + if err != nil { + return nil, infraerrors.Newf(http.StatusBadGateway, "GROK_OAUTH_CLIENT_INIT_FAILED", "create HTTP client: %v", err) + } + + clientID = strings.TrimSpace(clientID) + if clientID == "" { + clientID = xai.EffectiveClientID() + } + + formData := url.Values{} + formData.Set("grant_type", "refresh_token") + formData.Set("client_id", clientID) + formData.Set("refresh_token", refreshToken) + + var tokenResp xai.TokenResponse + resp, err := client.R(). + SetContext(ctx). + SetHeader("User-Agent", "sub2api-grok-oauth/1.0"). + SetFormDataFromValues(formData). + SetSuccessResult(&tokenResp). + Post(c.tokenURL) + if err != nil { + return nil, infraerrors.Newf(http.StatusBadGateway, "GROK_OAUTH_REQUEST_FAILED", "request failed: %v", err) + } + if !resp.IsSuccessState() { + return nil, grokOAuthStatusError("GROK_OAUTH_TOKEN_REFRESH_FAILED", "token refresh failed", resp) + } + return &tokenResp, nil +} + +func createGrokReqClient(proxyURL string) (*req.Client, error) { + return getSharedReqClient(reqClientOptions{ + ProxyURL: proxyURL, + Timeout: 60 * time.Second, + }) +} + +func grokOAuthStatusError(code, message string, resp *req.Response) error { + statusCode := http.StatusBadGateway + errorCode := code + upstreamStatus := 0 + if resp != nil && resp.StatusCode == http.StatusForbidden { + statusCode = http.StatusForbidden + errorCode = "GROK_OAUTH_ENTITLEMENT_DENIED" + } + body := "" + if resp != nil { + upstreamStatus = resp.StatusCode + body = logredact.RedactText(resp.String()) + } + return infraerrors.Newf(statusCode, errorCode, "%s: status %d, body: %s", message, upstreamStatus, body) +} diff --git a/backend/internal/repository/grok_oauth_client_test.go b/backend/internal/repository/grok_oauth_client_test.go new file mode 100644 index 0000000000..eabeb641c7 --- /dev/null +++ b/backend/internal/repository/grok_oauth_client_test.go @@ -0,0 +1,107 @@ +//go:build unit + +package repository + +import ( + "context" + "encoding/json" + "net/http" + "net/http/httptest" + "strings" + "testing" + + "github.com/Wei-Shaw/sub2api/internal/pkg/xai" + "github.com/stretchr/testify/require" +) + +func TestGrokOAuthClientExchangeAndRefreshUseFormFields(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + require.Equal(t, http.MethodPost, r.Method) + require.NoError(t, r.ParseForm()) + require.Equal(t, "client-id", r.Form.Get("client_id")) + + switch r.Form.Get("grant_type") { + case "authorization_code": + require.Equal(t, "auth-code", r.Form.Get("code")) + require.Equal(t, "http://127.0.0.1:56121/callback", r.Form.Get("redirect_uri")) + require.Equal(t, "verifier", r.Form.Get("code_verifier")) + require.Empty(t, r.Form.Get("code_challenge")) + require.Empty(t, r.Form.Get("code_challenge_method")) + _ = json.NewEncoder(w).Encode(map[string]any{ + "access_token": "exchange-access", + "refresh_token": "exchange-refresh", + "token_type": "Bearer", + "expires_in": 3600, + "scope": "openid api:access", + }) + case "refresh_token": + require.Equal(t, "refresh-token", r.Form.Get("refresh_token")) + _ = json.NewEncoder(w).Encode(map[string]any{ + "access_token": "refresh-access", + "refresh_token": "refresh-rotated", + "token_type": "Bearer", + "expires_in": 7200, + }) + default: + http.Error(w, "unexpected grant_type", http.StatusBadRequest) + } + })) + defer server.Close() + t.Setenv(xai.EnvTokenURL, server.URL) + + client := NewGrokOAuthClient() + + exchanged, err := client.ExchangeCode( + context.Background(), + "auth-code", + "verifier", + "http://127.0.0.1:56121/callback", + "", + "client-id", + ) + require.NoError(t, err) + require.Equal(t, "exchange-access", exchanged.AccessToken) + require.Equal(t, "exchange-refresh", exchanged.RefreshToken) + require.Equal(t, int64(3600), exchanged.ExpiresIn) + require.Equal(t, "openid api:access", exchanged.Scope) + + refreshed, err := client.RefreshToken(context.Background(), "refresh-token", "", "client-id") + require.NoError(t, err) + require.Equal(t, "refresh-access", refreshed.AccessToken) + require.Equal(t, "refresh-rotated", refreshed.RefreshToken) + require.Equal(t, int64(7200), refreshed.ExpiresIn) +} + +func TestGrokOAuthClientRefreshForbiddenClassifiesEntitlement(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + w.WriteHeader(http.StatusForbidden) + _, _ = w.Write([]byte(`{"error":"subscription required"}`)) + })) + defer server.Close() + t.Setenv(xai.EnvTokenURL, server.URL) + + client := NewGrokOAuthClient() + _, err := client.RefreshToken(context.Background(), "refresh-token", "", "client-id") + require.Error(t, err) + require.Contains(t, strings.ToUpper(err.Error()), "GROK_OAUTH_ENTITLEMENT_DENIED") +} + +func TestGrokOAuthClientStatusErrorRedactsSensitiveResponseBody(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + w.WriteHeader(http.StatusBadRequest) + _, _ = w.Write([]byte(`{"error":"invalid_grant","access_token":"access-secret","refresh_token":"refresh-secret","code_verifier":"verifier-secret"}`)) + })) + defer server.Close() + t.Setenv(xai.EnvTokenURL, server.URL) + + client := NewGrokOAuthClient() + _, err := client.RefreshToken(context.Background(), "refresh-secret", "", "client-id") + require.Error(t, err) + + errText := err.Error() + require.Contains(t, errText, "status 400") + require.Contains(t, errText, `\"refresh_token\":\"***\"`) + require.NotContains(t, errText, "access-secret") + require.NotContains(t, errText, "refresh-secret") + require.NotContains(t, errText, "verifier-secret") +} diff --git a/backend/internal/repository/simple_mode_default_groups.go b/backend/internal/repository/simple_mode_default_groups.go index 5630918400..e3786451a4 100644 --- a/backend/internal/repository/simple_mode_default_groups.go +++ b/backend/internal/repository/simple_mode_default_groups.go @@ -19,6 +19,7 @@ func ensureSimpleModeDefaultGroups(ctx context.Context, client *dbent.Client) er service.PlatformOpenAI: 1, service.PlatformGemini: 1, service.PlatformAntigravity: 2, + service.PlatformGrok: 1, } for platform, minCount := range requiredByPlatform { diff --git a/backend/internal/repository/wire.go b/backend/internal/repository/wire.go index 630c27ed47..37f8e9bd2f 100644 --- a/backend/internal/repository/wire.go +++ b/backend/internal/repository/wire.go @@ -141,6 +141,7 @@ var ProviderSet = wire.NewSet( NewClaudeOAuthClient, NewHTTPUpstream, NewOpenAIOAuthClient, + NewGrokOAuthClient, NewGeminiOAuthClient, NewGeminiCliCodeAssistClient, NewGeminiDriveClient, diff --git a/backend/internal/server/routes/admin.go b/backend/internal/server/routes/admin.go index 224300f12f..8728ea2c74 100644 --- a/backend/internal/server/routes/admin.go +++ b/backend/internal/server/routes/admin.go @@ -47,6 +47,9 @@ func RegisterAdminRoutes( // Antigravity OAuth registerAntigravityOAuthRoutes(admin, h) + // Grok OAuth + registerGrokOAuthRoutes(admin, h) + // 代理管理 registerProxyRoutes(admin, h) @@ -386,6 +389,20 @@ func registerAntigravityOAuthRoutes(admin *gin.RouterGroup, h *handler.Handlers) } } +func registerGrokOAuthRoutes(admin *gin.RouterGroup, h *handler.Handlers) { + grok := admin.Group("/grok") + { + grok.POST("/oauth/auth-url", h.Admin.GrokOAuth.GenerateAuthURL) + grok.POST("/oauth/exchange-code", h.Admin.GrokOAuth.ExchangeCode) + grok.POST("/oauth/refresh-token", h.Admin.GrokOAuth.RefreshToken) + grok.POST("/oauth/create-from-oauth", h.Admin.GrokOAuth.CreateAccountFromOAuth) + grok.POST("/accounts/:id/refresh", h.Admin.GrokOAuth.RefreshAccountToken) + grok.GET("/accounts/:id/quota", h.Admin.GrokOAuth.QueryQuota) + grok.POST("/accounts/:id/reset-quota", h.Admin.GrokOAuth.ResetQuota) + grok.GET("/runtime-sanity", h.Admin.GrokOAuth.RuntimeSanity) + } +} + func registerProxyRoutes(admin *gin.RouterGroup, h *handler.Handlers) { proxies := admin.Group("/proxies") { diff --git a/backend/internal/server/routes/gateway.go b/backend/internal/server/routes/gateway.go index b039a6ecd4..0d0c8852d9 100644 --- a/backend/internal/server/routes/gateway.go +++ b/backend/internal/server/routes/gateway.go @@ -31,6 +31,27 @@ func RegisterGatewayRoutes( requireGroupAnthropic := middleware.RequireGroupAssignment(settingService, middleware.AnthropicErrorWriter) requireGroupGoogle := middleware.RequireGroupAssignment(settingService, middleware.GoogleErrorWriter) + isOpenAIResponsesCompatibleGatewayPlatform := func(c *gin.Context) bool { + switch getGroupPlatform(c) { + case service.PlatformOpenAI, service.PlatformGrok: + return true + default: + return false + } + } + isOpenAIGatewayPlatform := func(c *gin.Context) bool { + return getGroupPlatform(c) == service.PlatformOpenAI + } + rejectGrokUnsupportedEndpoint := func(c *gin.Context, endpoint string) { + service.MarkOpsClientBusinessLimited(c, service.OpsClientBusinessLimitedReasonLocalFeatureGate) + c.JSON(http.StatusNotFound, gin.H{ + "error": gin.H{ + "type": "not_found_error", + "message": endpoint + " is not supported for Grok groups", + }, + }) + } + // API网关(Claude API兼容) gateway := r.Group("/v1") gateway.Use(bodyLimit) @@ -42,7 +63,11 @@ func RegisterGatewayRoutes( { // /v1/messages: auto-route based on group platform gateway.POST("/messages", func(c *gin.Context) { - if getGroupPlatform(c) == service.PlatformOpenAI { + if getGroupPlatform(c) == service.PlatformGrok { + rejectGrokUnsupportedEndpoint(c, "Messages API") + return + } + if isOpenAIGatewayPlatform(c) { h.OpenAIGateway.Messages(c) return } @@ -50,7 +75,7 @@ func RegisterGatewayRoutes( }) // /v1/messages/count_tokens: OpenAI groups get 404 gateway.POST("/messages/count_tokens", func(c *gin.Context) { - if getGroupPlatform(c) == service.PlatformOpenAI { + if isOpenAIResponsesCompatibleGatewayPlatform(c) { service.MarkOpsClientBusinessLimited(c, service.OpsClientBusinessLimitedReasonLocalFeatureGate) c.JSON(http.StatusNotFound, gin.H{ "type": "error", @@ -67,23 +92,33 @@ func RegisterGatewayRoutes( gateway.GET("/usage", h.Gateway.Usage) // OpenAI Responses API: auto-route based on group platform gateway.POST("/responses", func(c *gin.Context) { - if getGroupPlatform(c) == service.PlatformOpenAI { + if isOpenAIResponsesCompatibleGatewayPlatform(c) { h.OpenAIGateway.Responses(c) return } h.Gateway.Responses(c) }) gateway.POST("/responses/*subpath", func(c *gin.Context) { - if getGroupPlatform(c) == service.PlatformOpenAI { + if isOpenAIResponsesCompatibleGatewayPlatform(c) { h.OpenAIGateway.Responses(c) return } h.Gateway.Responses(c) }) - gateway.GET("/responses", h.OpenAIGateway.ResponsesWebSocket) + gateway.GET("/responses", func(c *gin.Context) { + if getGroupPlatform(c) == service.PlatformGrok { + rejectGrokUnsupportedEndpoint(c, "Responses WebSocket API") + return + } + h.OpenAIGateway.ResponsesWebSocket(c) + }) // OpenAI Chat Completions API: auto-route based on group platform gateway.POST("/chat/completions", func(c *gin.Context) { - if getGroupPlatform(c) == service.PlatformOpenAI { + if getGroupPlatform(c) == service.PlatformGrok { + rejectGrokUnsupportedEndpoint(c, "Chat Completions API") + return + } + if isOpenAIGatewayPlatform(c) { h.OpenAIGateway.ChatCompletions(c) return } @@ -147,7 +182,7 @@ func RegisterGatewayRoutes( // OpenAI Responses API(不带v1前缀的别名)— auto-route based on group platform responsesHandler := func(c *gin.Context) { - if getGroupPlatform(c) == service.PlatformOpenAI { + if isOpenAIResponsesCompatibleGatewayPlatform(c) { h.OpenAIGateway.Responses(c) return } @@ -155,17 +190,33 @@ func RegisterGatewayRoutes( } r.POST("/responses", bodyLimit, clientRequestID, opsErrorLogger, endpointNorm, gin.HandlerFunc(apiKeyAuth), requireGroupAnthropic, responsesHandler) r.POST("/responses/*subpath", bodyLimit, clientRequestID, opsErrorLogger, endpointNorm, gin.HandlerFunc(apiKeyAuth), requireGroupAnthropic, responsesHandler) - r.GET("/responses", bodyLimit, clientRequestID, opsErrorLogger, endpointNorm, gin.HandlerFunc(apiKeyAuth), requireGroupAnthropic, h.OpenAIGateway.ResponsesWebSocket) + r.GET("/responses", bodyLimit, clientRequestID, opsErrorLogger, endpointNorm, gin.HandlerFunc(apiKeyAuth), requireGroupAnthropic, func(c *gin.Context) { + if getGroupPlatform(c) == service.PlatformGrok { + rejectGrokUnsupportedEndpoint(c, "Responses WebSocket API") + return + } + h.OpenAIGateway.ResponsesWebSocket(c) + }) codexDirect := r.Group("/backend-api/codex") codexDirect.Use(bodyLimit, clientRequestID, opsErrorLogger, endpointNorm, gin.HandlerFunc(apiKeyAuth), requireGroupAnthropic) { codexDirect.POST("/responses", responsesHandler) codexDirect.POST("/responses/*subpath", responsesHandler) - codexDirect.GET("/responses", h.OpenAIGateway.ResponsesWebSocket) + codexDirect.GET("/responses", func(c *gin.Context) { + if getGroupPlatform(c) == service.PlatformGrok { + rejectGrokUnsupportedEndpoint(c, "Responses WebSocket API") + return + } + h.OpenAIGateway.ResponsesWebSocket(c) + }) } // OpenAI Chat Completions API(不带v1前缀的别名)— auto-route based on group platform r.POST("/chat/completions", bodyLimit, clientRequestID, opsErrorLogger, endpointNorm, gin.HandlerFunc(apiKeyAuth), requireGroupAnthropic, func(c *gin.Context) { - if getGroupPlatform(c) == service.PlatformOpenAI { + if getGroupPlatform(c) == service.PlatformGrok { + rejectGrokUnsupportedEndpoint(c, "Chat Completions API") + return + } + if isOpenAIGatewayPlatform(c) { h.OpenAIGateway.ChatCompletions(c) return } diff --git a/backend/internal/server/routes/gateway_test.go b/backend/internal/server/routes/gateway_test.go index 19ef568600..16cab29eaf 100644 --- a/backend/internal/server/routes/gateway_test.go +++ b/backend/internal/server/routes/gateway_test.go @@ -14,10 +14,15 @@ import ( "github.com/stretchr/testify/require" ) -func newGatewayRoutesTestRouter() *gin.Engine { +func newGatewayRoutesTestRouter(platform ...string) *gin.Engine { gin.SetMode(gin.TestMode) router := gin.New() + groupPlatform := service.PlatformOpenAI + if len(platform) > 0 && platform[0] != "" { + groupPlatform = platform[0] + } + RegisterGatewayRoutes( router, &handler.Handlers{ @@ -28,7 +33,7 @@ func newGatewayRoutesTestRouter() *gin.Engine { groupID := int64(1) c.Set(string(servermiddleware.ContextKeyAPIKey), &service.APIKey{ GroupID: &groupID, - Group: &service.Group{Platform: service.PlatformOpenAI}, + Group: &service.Group{Platform: groupPlatform}, }) c.Next() }), @@ -77,3 +82,40 @@ func TestGatewayRoutesOpenAIImagesPathsAreRegistered(t *testing.T) { require.NotEqual(t, http.StatusNotFound, w.Code, "path=%s should hit OpenAI images handler", path) } } + +func TestGatewayRoutesGrokOnlyAllowsResponsesHTTP(t *testing.T) { + router := newGatewayRoutesTestRouter(service.PlatformGrok) + + for _, tc := range []struct { + method string + path string + }{ + {http.MethodPost, "/v1/messages"}, + {http.MethodPost, "/v1/chat/completions"}, + {http.MethodPost, "/chat/completions"}, + {http.MethodGet, "/v1/responses"}, + {http.MethodGet, "/responses"}, + {http.MethodGet, "/backend-api/codex/responses"}, + } { + req := httptest.NewRequest(tc.method, tc.path, strings.NewReader(`{"model":"grok"}`)) + req.Header.Set("Content-Type", "application/json") + w := httptest.NewRecorder() + + router.ServeHTTP(w, req) + require.Equal(t, http.StatusNotFound, w.Code, "method=%s path=%s", tc.method, tc.path) + require.Contains(t, w.Body.String(), "not supported for Grok groups") + } + + for _, path := range []string{ + "/v1/responses", + "/responses", + "/backend-api/codex/responses", + } { + req := httptest.NewRequest(http.MethodPost, path, strings.NewReader(`{"model":"grok","input":"hi"}`)) + req.Header.Set("Content-Type", "application/json") + w := httptest.NewRecorder() + + router.ServeHTTP(w, req) + require.NotEqual(t, http.StatusNotFound, w.Code, "path=%s should still reach Responses handler", path) + } +} diff --git a/backend/internal/service/account.go b/backend/internal/service/account.go index 1f86be69b1..5c3ac71ea2 100644 --- a/backend/internal/service/account.go +++ b/backend/internal/service/account.go @@ -14,6 +14,7 @@ import ( "github.com/Wei-Shaw/sub2api/internal/config" "github.com/Wei-Shaw/sub2api/internal/domain" + "github.com/Wei-Shaw/sub2api/internal/pkg/xai" ) type Account struct { @@ -190,6 +191,18 @@ func (a *Account) IsGemini() bool { return a.Platform == PlatformGemini } +func (a *Account) IsGrok() bool { + return a.Platform == PlatformGrok +} + +func (a *Account) IsGrokOAuth() bool { + return a.IsGrok() && a.Type == AccountTypeOAuth +} + +func (a *Account) IsOpenAICompatible() bool { + return a != nil && (a.Platform == PlatformOpenAI || a.Platform == PlatformGrok) +} + func (a *Account) GeminiOAuthType() string { if a.Platform != PlatformGemini || a.Type != AccountTypeOAuth { return "" @@ -508,6 +521,9 @@ func (a *Account) resolveModelMapping(rawMapping map[string]any) map[string]stri if a.Platform == domain.PlatformAntigravity { return domain.DefaultAntigravityModelMapping } + if a.Platform == domain.PlatformGrok { + return xai.DefaultModelMapping() + } // Bedrock 默认映射由 forwardBedrock 统一处理(需配合 region prefix 调整) return nil } @@ -516,6 +532,9 @@ func (a *Account) resolveModelMapping(rawMapping map[string]any) map[string]stri if a.Platform == domain.PlatformAntigravity { return domain.DefaultAntigravityModelMapping } + if a.Platform == domain.PlatformGrok { + return xai.DefaultModelMapping() + } return nil } @@ -540,6 +559,9 @@ func (a *Account) resolveModelMapping(rawMapping map[string]any) map[string]stri if a.Platform == domain.PlatformAntigravity { return domain.DefaultAntigravityModelMapping } + if a.Platform == domain.PlatformGrok { + return xai.DefaultModelMapping() + } return nil } @@ -1114,6 +1136,31 @@ func (a *Account) GetOpenAIRefreshToken() string { return a.GetCredential("refresh_token") } +func (a *Account) GetGrokBaseURL() string { + if !a.IsGrok() { + return "" + } + baseURL := a.GetCredential("base_url") + if baseURL != "" { + return baseURL + } + return xai.DefaultBaseURL +} + +func (a *Account) GetGrokAccessToken() string { + if !a.IsGrok() { + return "" + } + return a.GetCredential("access_token") +} + +func (a *Account) GetGrokRefreshToken() string { + if !a.IsGrokOAuth() { + return "" + } + return a.GetCredential("refresh_token") +} + func (a *Account) GetOpenAIIDToken() string { if !a.IsOpenAIOAuth() { return "" @@ -1191,9 +1238,12 @@ func (a *Account) SupportsOpenAIEndpointCapability(capability OpenAIEndpointCapa if capability == "" { return true } - if !a.IsOpenAI() { + if !a.IsOpenAICompatible() { return false } + if a.IsGrok() { + return capability == OpenAIEndpointCapabilityChatCompletions + } switch capability { case OpenAIEndpointCapabilityChatCompletions: case OpenAIEndpointCapabilityEmbeddings: @@ -1259,6 +1309,9 @@ func (a *Account) openAIEndpointCapabilitySet() (map[string]bool, bool) { } func (a *Account) SupportsOpenAIImageCapability(capability OpenAIImagesCapability) bool { + if capability == "" { + return true + } if !a.IsOpenAI() { return false } diff --git a/backend/internal/service/account_service.go b/backend/internal/service/account_service.go index 41b5a2d97e..de3e9e7d5a 100644 --- a/backend/internal/service/account_service.go +++ b/backend/internal/service/account_service.go @@ -188,7 +188,7 @@ func (s *AccountService) Create(ctx context.Context, req CreateAccountRequest) ( if err != nil { return nil, err } - if g.RequireOAuthOnly && (g.Platform == PlatformOpenAI || g.Platform == PlatformAntigravity || g.Platform == PlatformAnthropic || g.Platform == PlatformGemini) { + if g.RequireOAuthOnly && (g.Platform == PlatformOpenAI || g.Platform == PlatformAntigravity || g.Platform == PlatformAnthropic || g.Platform == PlatformGemini || g.Platform == PlatformGrok) { return nil, fmt.Errorf("分组 [%s] 仅允许 OAuth 账号,apikey 类型账号无法加入", g.Name) } } @@ -304,7 +304,7 @@ func (s *AccountService) Update(ctx context.Context, id int64, req UpdateAccount if err != nil { return nil, err } - if g.RequireOAuthOnly && (g.Platform == PlatformOpenAI || g.Platform == PlatformAntigravity || g.Platform == PlatformAnthropic || g.Platform == PlatformGemini) { + if g.RequireOAuthOnly && (g.Platform == PlatformOpenAI || g.Platform == PlatformAntigravity || g.Platform == PlatformAnthropic || g.Platform == PlatformGemini || g.Platform == PlatformGrok) { return nil, fmt.Errorf("分组 [%s] 仅允许 OAuth 账号,apikey 类型账号无法加入", g.Name) } } @@ -427,6 +427,9 @@ func (s *AccountService) TestCredentials(ctx context.Context, id int64) error { case PlatformGemini: // TODO: 测试Gemini API凭证 return nil + case PlatformGrok: + // Grok OAuth credentials are validated via token exchange/refresh and request-path probes. + return nil default: return fmt.Errorf("unsupported platform: %s", account.Platform) } diff --git a/backend/internal/service/account_usage_service.go b/backend/internal/service/account_usage_service.go index 8092688d89..2963d996ac 100644 --- a/backend/internal/service/account_usage_service.go +++ b/backend/internal/service/account_usage_service.go @@ -19,6 +19,7 @@ import ( "github.com/Wei-Shaw/sub2api/internal/pkg/timezone" "github.com/Wei-Shaw/sub2api/internal/pkg/tlsfingerprint" "github.com/Wei-Shaw/sub2api/internal/pkg/usagestats" + "github.com/Wei-Shaw/sub2api/internal/pkg/xai" "golang.org/x/sync/errgroup" "golang.org/x/sync/singleflight" ) @@ -193,6 +194,17 @@ type UsageInfo struct { // Antigravity 多模型配额 AntigravityQuota map[string]*AntigravityModelQuota `json:"antigravity_quota,omitempty"` + // Grok / xAI 被动额度快照 + GrokRequestQuota *xai.QuotaWindow `json:"grok_request_quota,omitempty"` + GrokTokenQuota *xai.QuotaWindow `json:"grok_token_quota,omitempty"` + GrokRetryAfterSeconds *int `json:"grok_retry_after_seconds,omitempty"` + GrokEntitlementStatus string `json:"grok_entitlement_status,omitempty"` + GrokQuotaSnapshotState string `json:"grok_quota_snapshot_state,omitempty"` + GrokLastQuotaProbeAt string `json:"grok_last_quota_probe_at,omitempty"` + GrokLastHeadersSeenAt string `json:"grok_last_headers_seen_at,omitempty"` + GrokLastStatusCode int `json:"grok_last_status_code,omitempty"` + GrokLocalUsage *WindowStats `json:"grok_local_usage,omitempty"` + // Antigravity 账号级信息 SubscriptionTier string `json:"subscription_tier,omitempty"` // 归一化订阅等级: FREE/PRO/ULTRA/UNKNOWN SubscriptionTierRaw string `json:"subscription_tier_raw,omitempty"` // 上游原始订阅等级名称 @@ -263,6 +275,7 @@ type AccountUsageService struct { usageFetcher ClaudeUsageFetcher geminiQuotaService *GeminiQuotaService antigravityQuotaFetcher *AntigravityQuotaFetcher + grokQuotaFetcher *GrokQuotaFetcher cache *UsageCache identityCache IdentityCache tlsFPProfileService *TLSFingerprintProfileService @@ -275,6 +288,7 @@ func NewAccountUsageService( usageFetcher ClaudeUsageFetcher, geminiQuotaService *GeminiQuotaService, antigravityQuotaFetcher *AntigravityQuotaFetcher, + grokQuotaFetcher *GrokQuotaFetcher, cache *UsageCache, identityCache IdentityCache, tlsFPProfileService *TLSFingerprintProfileService, @@ -285,6 +299,7 @@ func NewAccountUsageService( usageFetcher: usageFetcher, geminiQuotaService: geminiQuotaService, antigravityQuotaFetcher: antigravityQuotaFetcher, + grokQuotaFetcher: grokQuotaFetcher, cache: cache, identityCache: identityCache, tlsFPProfileService: tlsFPProfileService, @@ -328,6 +343,14 @@ func (s *AccountUsageService) GetUsage(ctx context.Context, accountID int64, for return usage, err } + if account.Platform == PlatformGrok { + usage, err := s.getGrokUsage(ctx, account) + if err == nil { + s.tryClearRecoverableAccountError(ctx, account) + } + return usage, err + } + // 只有oauth类型账号可以通过API获取usage(有profile scope) if account.CanGetUsage() { var apiResp *ClaudeUsageResponse @@ -837,6 +860,30 @@ func (s *AccountUsageService) getAntigravityUsage(ctx context.Context, account * return usage, nil } +func (s *AccountUsageService) getGrokUsage(ctx context.Context, account *Account) (*UsageInfo, error) { + if s.grokQuotaFetcher == nil { + now := time.Now() + return &UsageInfo{UpdatedAt: &now}, nil + } + usage := s.grokQuotaFetcher.BuildUsageInfo(account) + if usage.GrokQuotaSnapshotState == "" { + if usage.ErrorCode == "quota_unknown" { + usage.GrokQuotaSnapshotState = "unknown_until_first_response" + } else { + usage.GrokQuotaSnapshotState = "observed" + } + } + + if s.usageLogRepo != nil && account != nil { + if stats, err := s.usageLogRepo.GetAccountTodayStats(ctx, account.ID); err == nil && stats != nil { + usage.GrokLocalUsage = windowStatsFromAccountStats(stats) + } + } + + enrichUsageWithAccountError(usage, account) + return usage, nil +} + // recalcAntigravityRemainingSeconds 重新计算 Antigravity UsageInfo 中各窗口的 RemainingSeconds // 用于从缓存取出时更新倒计时,避免返回过时的剩余秒数 func recalcAntigravityRemainingSeconds(info *UsageInfo) { 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 5a7c14f00a..7a5637dc75 100644 --- a/backend/internal/service/admin_service.go +++ b/backend/internal/service/admin_service.go @@ -25,6 +25,7 @@ import ( "github.com/Wei-Shaw/sub2api/internal/pkg/logger" "github.com/Wei-Shaw/sub2api/internal/pkg/openai" "github.com/Wei-Shaw/sub2api/internal/pkg/pagination" + "github.com/Wei-Shaw/sub2api/internal/pkg/xai" "github.com/Wei-Shaw/sub2api/internal/util/httputil" ) @@ -1780,6 +1781,8 @@ func defaultModelsListCandidateIDs(platform string) []string { ids = append(ids, model.ID) } return ids + case PlatformGrok: + return xai.DefaultModelIDs() default: ids := make([]string, 0, len(claude.DefaultModels)) for _, model := range claude.DefaultModels { @@ -1913,7 +1916,7 @@ func (s *adminServiceImpl) CreateGroup(ctx context.Context, input *CreateGroupIn } // require_oauth_only: 过滤掉 apikey 类型账号 - if group.RequireOAuthOnly && (group.Platform == PlatformOpenAI || group.Platform == PlatformAntigravity || group.Platform == PlatformAnthropic || group.Platform == PlatformGemini) && len(accountIDsToCopy) > 0 { + if group.RequireOAuthOnly && (group.Platform == PlatformOpenAI || group.Platform == PlatformAntigravity || group.Platform == PlatformAnthropic || group.Platform == PlatformGemini || group.Platform == PlatformGrok) && len(accountIDsToCopy) > 0 { accounts, err := s.accountRepo.GetByIDs(ctx, accountIDsToCopy) if err != nil { return nil, fmt.Errorf("failed to fetch accounts for oauth filter: %w", err) @@ -2208,7 +2211,7 @@ func (s *adminServiceImpl) UpdateGroup(ctx context.Context, id int64, input *Upd } // require_oauth_only: 过滤掉 apikey 类型账号 - if group.RequireOAuthOnly && (group.Platform == PlatformOpenAI || group.Platform == PlatformAntigravity || group.Platform == PlatformAnthropic || group.Platform == PlatformGemini) && len(accountIDsToCopy) > 0 { + if group.RequireOAuthOnly && (group.Platform == PlatformOpenAI || group.Platform == PlatformAntigravity || group.Platform == PlatformAnthropic || group.Platform == PlatformGemini || group.Platform == PlatformGrok) && len(accountIDsToCopy) > 0 { accounts, err := s.accountRepo.GetByIDs(ctx, accountIDsToCopy) if err != nil { return nil, fmt.Errorf("failed to fetch accounts for oauth filter: %w", err) @@ -2569,6 +2572,18 @@ 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 { + if concurrency <= 0 { + return 1 + } + if concurrency > 1 && !xai.AllowUnsafeHighConcurrency() { + return 1 + } + } + return concurrency +} + func (s *adminServiceImpl) CreateAccount(ctx context.Context, input *CreateAccountInput) (*Account, error) { // 绑定分组 groupIDs := input.GroupIDs @@ -2601,7 +2616,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, @@ -2734,7 +2749,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 { diff --git a/backend/internal/service/billing_service.go b/backend/internal/service/billing_service.go index 74055a8151..b8bb555a77 100644 --- a/backend/internal/service/billing_service.go +++ b/backend/internal/service/billing_service.go @@ -499,6 +499,22 @@ func (s *BillingService) initFallbackPricing() { OutputPricePerToken: 0, SupportsCacheBreakdown: false, } + + // xAI Grok 4.3 (official docs: $1.25 input / $2.50 output per MTok) + s.fallbackPrices["grok-4.3"] = &ModelPricing{ + InputPricePerToken: 1.25e-6, + OutputPricePerToken: 2.5e-6, + CacheReadPricePerToken: 0, + SupportsCacheBreakdown: false, + LongContextInputThreshold: 1000000, + LongContextInputMultiplier: 1, + } + // xAI Grok Build 0.1 (official docs: $1 input / $2 output per MTok) + s.fallbackPrices["grok-build-0.1"] = &ModelPricing{ + InputPricePerToken: 1e-6, + OutputPricePerToken: 2e-6, + SupportsCacheBreakdown: false, + } } // getFallbackPricing 根据模型系列获取回退价格 @@ -659,6 +675,13 @@ func (s *BillingService) getFallbackPricing(model string) *ModelPricing { } } + switch modelLower { + case "grok", "grok-latest", "grok-4.3": + return s.fallbackPrices["grok-4.3"] + case "grok-build", "grok-build-0.1": + return s.fallbackPrices["grok-build-0.1"] + } + return nil } diff --git a/backend/internal/service/domain_constants.go b/backend/internal/service/domain_constants.go index 3908c00abd..7dd5ec6323 100644 --- a/backend/internal/service/domain_constants.go +++ b/backend/internal/service/domain_constants.go @@ -41,6 +41,7 @@ const ( PlatformOpenAI = domain.PlatformOpenAI PlatformGemini = domain.PlatformGemini PlatformAntigravity = domain.PlatformAntigravity + PlatformGrok = domain.PlatformGrok ) // AllowedQuotaPlatforms 是允许设置 user × platform quota 的平台列表(单一权威来源)。 @@ -51,6 +52,7 @@ var AllowedQuotaPlatforms = []string{ PlatformOpenAI, PlatformGemini, PlatformAntigravity, + PlatformGrok, } // IsAllowedQuotaPlatform 报告 s 是否为合法的 quota platform 标识。 diff --git a/backend/internal/service/grok_oauth_service.go b/backend/internal/service/grok_oauth_service.go new file mode 100644 index 0000000000..acad39985f --- /dev/null +++ b/backend/internal/service/grok_oauth_service.go @@ -0,0 +1,312 @@ +package service + +import ( + "context" + "crypto/subtle" + "encoding/base64" + "encoding/json" + "net/http" + "strings" + "time" + + infraerrors "github.com/Wei-Shaw/sub2api/internal/pkg/errors" + "github.com/Wei-Shaw/sub2api/internal/pkg/xai" +) + +const grokDefaultAccessTokenTTL = 6 * time.Hour + +type GrokOAuthService struct { + sessionStore *xai.SessionStore + proxyRepo ProxyRepository + oauthClient GrokOAuthClient +} + +func NewGrokOAuthService(proxyRepo ProxyRepository, oauthClient GrokOAuthClient) *GrokOAuthService { + return &GrokOAuthService{ + sessionStore: xai.NewSessionStore(), + proxyRepo: proxyRepo, + oauthClient: oauthClient, + } +} + +type GrokAuthURLResult struct { + AuthURL string `json:"auth_url"` + SessionID string `json:"session_id"` + State string `json:"state"` +} + +func (s *GrokOAuthService) GenerateAuthURL(ctx context.Context, proxyID *int64, redirectURI string) (*GrokAuthURLResult, error) { + state, err := xai.GenerateState() + if err != nil { + return nil, infraerrors.Newf(http.StatusInternalServerError, "GROK_OAUTH_STATE_FAILED", "failed to generate state: %v", err) + } + nonce, err := xai.GenerateNonce() + if err != nil { + return nil, infraerrors.Newf(http.StatusInternalServerError, "GROK_OAUTH_NONCE_FAILED", "failed to generate nonce: %v", err) + } + codeVerifier, err := xai.GenerateCodeVerifier() + if err != nil { + return nil, infraerrors.Newf(http.StatusInternalServerError, "GROK_OAUTH_VERIFIER_FAILED", "failed to generate code verifier: %v", err) + } + sessionID, err := xai.GenerateSessionID() + if err != nil { + return nil, infraerrors.Newf(http.StatusInternalServerError, "GROK_OAUTH_SESSION_FAILED", "failed to generate session ID: %v", err) + } + + proxyURL, err := s.proxyURL(ctx, proxyID) + if err != nil { + return nil, err + } + redirectURI = xai.EffectiveRedirectURI(redirectURI) + codeChallenge := xai.GenerateCodeChallenge(codeVerifier) + + authURL, err := xai.BuildAuthorizationURL(state, codeChallenge, redirectURI, nonce) + if err != nil { + return nil, infraerrors.Newf(http.StatusBadRequest, "GROK_OAUTH_INVALID_AUTHORIZE_URL", "%v", err) + } + + s.sessionStore.Set(sessionID, &xai.OAuthSession{ + State: state, + CodeVerifier: codeVerifier, + CodeChallenge: codeChallenge, + ClientID: xai.EffectiveClientID(), + Scope: xai.EffectiveScope(), + ProxyURL: proxyURL, + RedirectURI: redirectURI, + CreatedAt: time.Now(), + }) + + return &GrokAuthURLResult{ + AuthURL: authURL, + SessionID: sessionID, + State: state, + }, nil +} + +type GrokExchangeCodeInput struct { + SessionID string + Code string + State string + RedirectURI string + ProxyID *int64 +} + +type GrokTokenInfo struct { + AccessToken string `json:"access_token"` + RefreshToken string `json:"refresh_token,omitempty"` + IDToken string `json:"id_token,omitempty"` + TokenType string `json:"token_type,omitempty"` + ExpiresIn int64 `json:"expires_in"` + ExpiresAt int64 `json:"expires_at"` + ClientID string `json:"client_id,omitempty"` + Scope string `json:"scope,omitempty"` + Email string `json:"email,omitempty"` + SubscriptionTier string `json:"subscription_tier,omitempty"` + EntitlementStatus string `json:"entitlement_status,omitempty"` +} + +func (s *GrokOAuthService) ExchangeCode(ctx context.Context, input *GrokExchangeCodeInput) (*GrokTokenInfo, error) { + if input == nil { + return nil, infraerrors.New(http.StatusBadRequest, "GROK_OAUTH_INVALID_INPUT", "input is required") + } + session, ok := s.sessionStore.Get(input.SessionID) + 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) + if code == "" { + return nil, infraerrors.New(http.StatusBadRequest, "GROK_OAUTH_CODE_REQUIRED", "authorization code is required") + } + state := strings.TrimSpace(input.State) + 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") + } + + proxyURL := session.ProxyURL + if input.ProxyID != nil { + var err error + proxyURL, err = s.proxyURL(ctx, input.ProxyID) + if err != nil { + return nil, err + } + } + redirectURI := session.RedirectURI + if strings.TrimSpace(input.RedirectURI) != "" { + redirectURI = input.RedirectURI + } + + tokenResp, err := s.oauthClient.ExchangeCode(ctx, code, session.CodeVerifier, redirectURI, proxyURL, session.ClientID) + if err != nil { + return nil, err + } + return s.tokenInfoFromResponse(tokenResp, session.ClientID, nil), nil +} + +func (s *GrokOAuthService) RefreshToken(ctx context.Context, refreshToken, proxyURL, clientID string) (*GrokTokenInfo, error) { + refreshToken = strings.TrimSpace(refreshToken) + if refreshToken == "" { + return nil, infraerrors.New(http.StatusBadRequest, "GROK_OAUTH_NO_REFRESH_TOKEN", "refresh_token is required") + } + tokenResp, err := s.oauthClient.RefreshToken(ctx, refreshToken, proxyURL, clientID) + if err != nil { + return nil, err + } + tokenInfo := s.tokenInfoFromResponse(tokenResp, clientID, nil) + if tokenInfo.RefreshToken == "" { + tokenInfo.RefreshToken = refreshToken + } + return tokenInfo, nil +} + +func (s *GrokOAuthService) ValidateRefreshToken(ctx context.Context, refreshToken string, proxyID *int64) (*GrokTokenInfo, error) { + proxyURL, err := s.proxyURL(ctx, proxyID) + if err != nil { + return nil, err + } + return s.RefreshToken(ctx, refreshToken, proxyURL, xai.EffectiveClientID()) +} + +func (s *GrokOAuthService) RefreshAccountToken(ctx context.Context, account *Account) (*GrokTokenInfo, error) { + if account == nil || account.Platform != PlatformGrok { + return nil, infraerrors.New(http.StatusBadRequest, "GROK_OAUTH_INVALID_ACCOUNT", "account is not a Grok account") + } + if account.Type != AccountTypeOAuth { + return nil, infraerrors.New(http.StatusBadRequest, "GROK_OAUTH_INVALID_ACCOUNT_TYPE", "account is not an OAuth account") + } + + proxyURL, err := s.proxyURL(ctx, account.ProxyID) + if err != nil { + return nil, err + } + refreshToken := account.GetCredential("refresh_token") + if strings.TrimSpace(refreshToken) == "" { + return nil, infraerrors.New(http.StatusBadRequest, "GROK_OAUTH_NO_REFRESH_TOKEN", "no refresh token available") + } + + clientID := account.GetCredential("client_id") + tokenInfo, err := s.RefreshToken(ctx, refreshToken, proxyURL, clientID) + if err != nil { + return nil, err + } + tokenInfo.SubscriptionTier = account.GetCredential("subscription_tier") + tokenInfo.EntitlementStatus = account.GetCredential("entitlement_status") + return tokenInfo, nil +} + +func (s *GrokOAuthService) BuildAccountCredentials(tokenInfo *GrokTokenInfo) map[string]any { + if tokenInfo == nil { + return nil + } + expiresAt := time.Unix(tokenInfo.ExpiresAt, 0).UTC().Format(time.RFC3339) + creds := map[string]any{ + "access_token": tokenInfo.AccessToken, + "expires_at": expiresAt, + } + if tokenInfo.RefreshToken != "" { + creds["refresh_token"] = tokenInfo.RefreshToken + } + if tokenInfo.TokenType != "" { + creds["token_type"] = tokenInfo.TokenType + } + if tokenInfo.IDToken != "" { + creds["id_token"] = tokenInfo.IDToken + } + if tokenInfo.ClientID != "" { + creds["client_id"] = tokenInfo.ClientID + } + if tokenInfo.Scope != "" { + creds["scope"] = tokenInfo.Scope + } + if tokenInfo.Email != "" { + creds["email"] = tokenInfo.Email + } + if tokenInfo.SubscriptionTier != "" { + creds["subscription_tier"] = tokenInfo.SubscriptionTier + } + if tokenInfo.EntitlementStatus != "" { + creds["entitlement_status"] = tokenInfo.EntitlementStatus + } + creds["base_url"] = xai.DefaultBaseURL + return creds +} + +func (s *GrokOAuthService) Stop() { + s.sessionStore.Stop() +} + +func (s *GrokOAuthService) tokenInfoFromResponse(tokenResp *xai.TokenResponse, clientID string, existing map[string]any) *GrokTokenInfo { + now := time.Now() + expiresIn := tokenResp.ExpiresIn + if expiresIn <= 0 { + expiresIn = int64(grokDefaultAccessTokenTTL.Seconds()) + } + info := &GrokTokenInfo{ + AccessToken: tokenResp.AccessToken, + RefreshToken: tokenResp.RefreshToken, + IDToken: tokenResp.IDToken, + TokenType: tokenResp.TokenType, + ExpiresIn: expiresIn, + ExpiresAt: now.Add(time.Duration(expiresIn) * time.Second).Unix(), + ClientID: strings.TrimSpace(clientID), + Scope: tokenResp.Scope, + } + if info.ClientID == "" { + info.ClientID = xai.EffectiveClientID() + } + if info.TokenType == "" { + info.TokenType = "Bearer" + } + if email := parseJWTEmailClaim(tokenResp.IDToken); email != "" { + info.Email = email + } + if info.Email == "" && existing != nil { + if email, _ := existing["email"].(string); email != "" { + info.Email = email + } + } + return info +} + +func (s *GrokOAuthService) proxyURL(ctx context.Context, proxyID *int64) (string, error) { + if proxyID == nil { + return "", nil + } + if s.proxyRepo == nil { + return "", infraerrors.New(http.StatusBadRequest, "GROK_OAUTH_PROXY_NOT_AVAILABLE", "proxy repository is not available") + } + proxy, err := s.proxyRepo.GetByID(ctx, *proxyID) + if err != nil { + return "", infraerrors.Newf(http.StatusBadRequest, "GROK_OAUTH_PROXY_NOT_FOUND", "proxy not found: %v", err) + } + if proxy == nil { + return "", nil + } + return proxy.URL(), nil +} + +func parseJWTEmailClaim(token string) string { + parts := strings.Split(token, ".") + if len(parts) < 2 { + return "" + } + payload, err := base64.RawURLEncoding.DecodeString(parts[1]) + if err != nil { + return "" + } + var claims struct { + Email string `json:"email"` + } + if err := json.Unmarshal(payload, &claims); err != nil { + return "" + } + return strings.TrimSpace(claims.Email) +} diff --git a/backend/internal/service/grok_oauth_service_test.go b/backend/internal/service/grok_oauth_service_test.go new file mode 100644 index 0000000000..d0caa5e527 --- /dev/null +++ b/backend/internal/service/grok_oauth_service_test.go @@ -0,0 +1,68 @@ +//go:build unit + +package service + +import ( + "context" + "testing" + + "github.com/Wei-Shaw/sub2api/internal/pkg/xai" + "github.com/stretchr/testify/require" +) + +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 +} + +func (s *grokOAuthClientStub) RefreshToken(context.Context, string, string, string) (*xai.TokenResponse, error) { + return s.refreshResponse, nil +} + +func TestGrokOAuthServiceRefreshTokenPreservesOriginalRefreshTokenWhenNotRotated(t *testing.T) { + svc := NewGrokOAuthService(nil, &grokOAuthClientStub{ + refreshResponse: &xai.TokenResponse{ + AccessToken: "new-access-token", + TokenType: "Bearer", + ExpiresIn: 3600, + }, + }) + defer svc.Stop() + + info, err := svc.RefreshToken(context.Background(), "original-refresh-token", "", "client-id") + require.NoError(t, err) + require.Equal(t, "new-access-token", info.AccessToken) + 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) +} diff --git a/backend/internal/service/grok_quota_fetcher.go b/backend/internal/service/grok_quota_fetcher.go new file mode 100644 index 0000000000..0939b78e20 --- /dev/null +++ b/backend/internal/service/grok_quota_fetcher.go @@ -0,0 +1,109 @@ +package service + +import ( + "encoding/json" + "fmt" + "time" + + "github.com/Wei-Shaw/sub2api/internal/pkg/xai" +) + +const grokQuotaSnapshotExtraKey = "grok_usage_snapshot" + +type GrokQuotaFetcher struct{} + +func NewGrokQuotaFetcher() *GrokQuotaFetcher { + return &GrokQuotaFetcher{} +} + +func (f *GrokQuotaFetcher) BuildUsageInfo(account *Account) *UsageInfo { + now := time.Now() + usage := &UsageInfo{ + Source: "passive", + UpdatedAt: &now, + } + if account == nil { + usage.ErrorCode = "quota_unknown" + usage.Error = "Grok quota is unknown until the first upstream response includes xAI rate-limit headers" + return usage + } + + snapshot, err := grokQuotaSnapshotFromExtra(account.Extra) + if err != nil || snapshot == nil { + usage.ErrorCode = "quota_unknown" + usage.Error = "Grok quota is unknown until the first upstream response includes xAI rate-limit headers" + return usage + } + + if parsedAt, err := time.Parse(time.RFC3339, snapshot.UpdatedAt); err == nil { + usage.UpdatedAt = &parsedAt + } + usage.GrokRequestQuota = snapshot.Requests + usage.GrokTokenQuota = snapshot.Tokens + usage.GrokRetryAfterSeconds = snapshot.RetryAfterSeconds + usage.SubscriptionTier = snapshot.SubscriptionTier + usage.SubscriptionTierRaw = snapshot.SubscriptionTier + usage.GrokEntitlementStatus = snapshot.EntitlementStatus + usage.GrokLastQuotaProbeAt = snapshot.LastProbeAt + usage.GrokLastHeadersSeenAt = snapshot.LastHeadersSeenAt + usage.GrokLastStatusCode = snapshot.StatusCode + if snapshot.HasObservedHeaders() { + usage.GrokQuotaSnapshotState = "observed" + } else { + usage.GrokQuotaSnapshotState = "no_headers" + usage.ErrorCode = "quota_unknown" + usage.Error = "No xAI quota headers observed on the latest Grok probe" + } + + switch snapshot.StatusCode { + case 401: + usage.NeedsReauth = true + usage.ErrorCode = "unauthenticated" + case 403: + usage.IsForbidden = true + usage.ForbiddenType = "forbidden" + usage.ErrorCode = "forbidden" + if usage.GrokEntitlementStatus == "" { + usage.GrokEntitlementStatus = "forbidden" + } + case 429: + usage.ErrorCode = "rate_limited" + } + return usage +} + +func grokQuotaSnapshotFromExtra(extra map[string]any) (*xai.QuotaSnapshot, error) { + if extra == nil { + return nil, nil + } + raw, ok := extra[grokQuotaSnapshotExtraKey] + if !ok || raw == nil { + return nil, nil + } + switch snapshot := raw.(type) { + case *xai.QuotaSnapshot: + return snapshot, nil + case xai.QuotaSnapshot: + return &snapshot, nil + case map[string]any: + data, err := json.Marshal(snapshot) + if err != nil { + return nil, err + } + var out xai.QuotaSnapshot + if err := json.Unmarshal(data, &out); err != nil { + return nil, err + } + return &out, nil + default: + data, err := json.Marshal(raw) + if err != nil { + return nil, fmt.Errorf("marshal grok quota snapshot: %w", err) + } + var out xai.QuotaSnapshot + if err := json.Unmarshal(data, &out); err != nil { + return nil, err + } + return &out, nil + } +} diff --git a/backend/internal/service/grok_quota_fetcher_test.go b/backend/internal/service/grok_quota_fetcher_test.go new file mode 100644 index 0000000000..d2d9c14993 --- /dev/null +++ b/backend/internal/service/grok_quota_fetcher_test.go @@ -0,0 +1,136 @@ +//go:build unit + +package service + +import ( + "net/http" + "testing" + "time" + + "github.com/Wei-Shaw/sub2api/internal/pkg/xai" + "github.com/stretchr/testify/require" +) + +func grokInt64PtrForTest(v int64) *int64 { return &v } +func grokIntPtrForTest(v int) *int { return &v } + +func TestGrokQuotaFetcherBuildUsageInfoUnknownUntilFirstSnapshot(t *testing.T) { + t.Parallel() + + usage := NewGrokQuotaFetcher().BuildUsageInfo(&Account{Platform: PlatformGrok, Type: AccountTypeOAuth}) + require.Equal(t, "passive", usage.Source) + require.Equal(t, "quota_unknown", usage.ErrorCode) + require.Contains(t, usage.Error, "unknown until the first upstream response") +} + +func TestGrokQuotaFetcherBuildUsageInfoFromSnapshot(t *testing.T) { + t.Parallel() + + updatedAt := "2030-01-01T00:00:00Z" + account := &Account{ + Platform: PlatformGrok, + Type: AccountTypeOAuth, + Extra: map[string]any{ + grokQuotaSnapshotExtraKey: &xai.QuotaSnapshot{ + Requests: &xai.QuotaWindow{ + Limit: grokInt64PtrForTest(100), + Remaining: grokInt64PtrForTest(12), + ResetAt: updatedAt, + }, + Tokens: &xai.QuotaWindow{ + Limit: grokInt64PtrForTest(1000), + Remaining: grokInt64PtrForTest(900), + }, + RetryAfterSeconds: grokIntPtrForTest(30), + SubscriptionTier: "supergrok", + EntitlementStatus: "active", + StatusCode: http.StatusTooManyRequests, + LastProbeAt: updatedAt, + LastHeadersSeenAt: updatedAt, + UpdatedAt: updatedAt, + }, + }, + } + + usage := NewGrokQuotaFetcher().BuildUsageInfo(account) + require.Equal(t, "passive", usage.Source) + require.Equal(t, "rate_limited", usage.ErrorCode) + require.Equal(t, "observed", usage.GrokQuotaSnapshotState) + require.Equal(t, "supergrok", usage.SubscriptionTier) + require.Equal(t, "active", usage.GrokEntitlementStatus) + require.Equal(t, int64(100), *usage.GrokRequestQuota.Limit) + require.Equal(t, int64(12), *usage.GrokRequestQuota.Remaining) + require.Equal(t, 30, *usage.GrokRetryAfterSeconds) + require.NotNil(t, usage.UpdatedAt) + require.Equal(t, updatedAt, usage.GrokLastQuotaProbeAt) + require.Equal(t, updatedAt, usage.GrokLastHeadersSeenAt) + require.Equal(t, http.StatusTooManyRequests, usage.GrokLastStatusCode) + require.True(t, usage.UpdatedAt.Equal(time.Date(2030, 1, 1, 0, 0, 0, 0, time.UTC))) +} + +func TestGrokQuotaFetcherBuildUsageInfoFromNoHeadersProbe(t *testing.T) { + t.Parallel() + + probedAt := "2030-01-01T00:00:00Z" + account := &Account{ + Platform: PlatformGrok, + Type: AccountTypeOAuth, + Extra: map[string]any{ + grokQuotaSnapshotExtraKey: xai.QuotaSnapshot{ + StatusCode: http.StatusOK, + HeadersObserved: false, + ObservationSource: "active_probe", + LastProbeAt: probedAt, + UpdatedAt: probedAt, + }, + }, + } + + usage := NewGrokQuotaFetcher().BuildUsageInfo(account) + require.Equal(t, "quota_unknown", usage.ErrorCode) + require.Equal(t, "no_headers", usage.GrokQuotaSnapshotState) + require.Contains(t, usage.Error, "No xAI quota headers observed") + require.Equal(t, probedAt, usage.GrokLastQuotaProbeAt) + require.Empty(t, usage.GrokLastHeadersSeenAt) + require.Equal(t, http.StatusOK, usage.GrokLastStatusCode) + require.Nil(t, usage.GrokRequestQuota) + require.Nil(t, usage.GrokTokenQuota) +} + +func TestGrokQuotaFetcherClassifiesForbiddenAndReauth(t *testing.T) { + t.Parallel() + + tests := []struct { + name string + statusCode int + wantReauth bool + wantForbid bool + wantCode string + wantEntitle string + }{ + {name: "reauth", statusCode: http.StatusUnauthorized, wantReauth: true, wantCode: "unauthenticated"}, + {name: "forbidden", statusCode: http.StatusForbidden, wantForbid: true, wantCode: "forbidden", wantEntitle: "forbidden"}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + t.Parallel() + account := &Account{ + Platform: PlatformGrok, + Type: AccountTypeOAuth, + Extra: map[string]any{ + grokQuotaSnapshotExtraKey: xai.QuotaSnapshot{ + StatusCode: tt.statusCode, + HeadersObserved: true, + UpdatedAt: time.Now().UTC().Format(time.RFC3339), + }, + }, + } + usage := NewGrokQuotaFetcher().BuildUsageInfo(account) + require.Equal(t, tt.wantReauth, usage.NeedsReauth) + require.Equal(t, tt.wantForbid, usage.IsForbidden) + require.Equal(t, tt.wantCode, usage.ErrorCode) + require.Equal(t, tt.wantEntitle, usage.GrokEntitlementStatus) + }) + } +} diff --git a/backend/internal/service/grok_quota_service.go b/backend/internal/service/grok_quota_service.go new file mode 100644 index 0000000000..19b1a01a0a --- /dev/null +++ b/backend/internal/service/grok_quota_service.go @@ -0,0 +1,198 @@ +package service + +import ( + "bytes" + "context" + "encoding/json" + "io" + "log/slog" + "net/http" + "strings" + "time" + + infraerrors "github.com/Wei-Shaw/sub2api/internal/pkg/errors" + "github.com/Wei-Shaw/sub2api/internal/pkg/xai" +) + +const ( + grokQuotaUpstreamTimeout = 20 * time.Second + grokQuotaProbeInput = "." + grokQuotaDefaultModel = "grok-4.3" +) + +type GrokQuotaProbeResult struct { + Source string `json:"source"` + Snapshot *xai.QuotaSnapshot `json:"snapshot,omitempty"` + StatusCode int `json:"status_code,omitempty"` + HeadersObserved bool `json:"headers_observed"` + ResetSupported bool `json:"reset_supported"` + FetchedAt int64 `json:"fetched_at"` +} + +type GrokQuotaResetResult struct { + Supported bool `json:"supported"` + Code string `json:"code"` + Message string `json:"message"` +} + +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, + } +} + +func (s *GrokQuotaService) ProbeUsage(ctx context.Context, accountID int64) (*GrokQuotaProbeResult, error) { + account, token, proxyURL, err := s.prepareProbe(ctx, accountID) + if err != nil { + return nil, err + } + + body, err := buildGrokQuotaProbeBody(account) + if err != nil { + return nil, infraerrors.Newf(http.StatusBadRequest, "GROK_QUOTA_PROBE_BODY_ERROR", "failed to build probe body: %v", err) + } + targetURL, err := xai.BuildResponsesURL(account.GetGrokBaseURL()) + if err != nil { + return nil, infraerrors.Newf(http.StatusBadRequest, "GROK_QUOTA_BASE_URL_INVALID", "invalid Grok base_url: %v", err) + } + + callCtx, cancel := context.WithTimeout(ctx, grokQuotaUpstreamTimeout) + defer cancel() + req, err := http.NewRequestWithContext(callCtx, http.MethodPost, targetURL, bytes.NewReader(body)) + if err != nil { + return nil, infraerrors.Newf(http.StatusInternalServerError, "GROK_QUOTA_PROBE_REQUEST_BUILD_FAILED", "failed to build upstream request: %v", err) + } + req.Header.Set("Authorization", "Bearer "+token) + req.Header.Set("Content-Type", "application/json") + req.Header.Set("Accept", "application/json") + req.Header.Set("User-Agent", "sub2api-grok-quota-probe/1.0") + + resp, err := s.httpUpstream.Do(req, proxyURL, account.ID, maxInt(account.Concurrency, 1)) + if err != nil { + return nil, infraerrors.Newf(http.StatusBadGateway, "GROK_QUOTA_PROBE_REQUEST_FAILED", "upstream probe failed: %v", err) + } + defer func() { _ = resp.Body.Close() }() + + snapshot := xai.ObserveQuotaHeaders(resp.Header, resp.StatusCode, "active_probe") + _ = s.accountRepo.UpdateExtra(ctx, account.ID, map[string]any{ + grokQuotaSnapshotExtraKey: snapshot, + }) + + result := &GrokQuotaProbeResult{ + Source: "active_probe", + Snapshot: snapshot, + StatusCode: resp.StatusCode, + HeadersObserved: snapshot.HeadersObserved, + ResetSupported: false, + FetchedAt: time.Now().Unix(), + } + if resp.StatusCode == http.StatusTooManyRequests { + return result, nil + } + if resp.StatusCode >= 400 { + bodyBytes, _ := io.ReadAll(io.LimitReader(resp.Body, 240)) + bodyText := truncate(strings.TrimSpace(string(bodyBytes)), 240) + slog.Warn("grok_quota_probe_failed", "account_id", account.ID, "status", resp.StatusCode, "body", bodyText) + return nil, infraerrors.Newf(mapUpstreamStatus(resp.StatusCode), "GROK_QUOTA_PROBE_UPSTREAM_ERROR", "upstream returned %d: %s", resp.StatusCode, bodyText) + } + return result, nil +} + +func (s *GrokQuotaService) ResetQuota(ctx context.Context, accountID int64) (*GrokQuotaResetResult, error) { + if _, err := s.loadGrokOAuthAccount(ctx, accountID); err != nil { + return nil, err + } + return nil, infraerrors.New(http.StatusNotImplemented, "GROK_QUOTA_RESET_UNSUPPORTED", "xAI does not expose a Grok subscription quota reset endpoint for OAuth accounts") +} + +func (s *GrokQuotaService) prepareProbe(ctx context.Context, accountID int64) (*Account, string, string, error) { + if s == nil || s.tokenProvider == nil || s.httpUpstream == nil { + return nil, "", "", infraerrors.New(http.StatusInternalServerError, "GROK_QUOTA_NOT_CONFIGURED", "grok quota service is not configured") + } + account, err := s.loadGrokOAuthAccount(ctx, accountID) + if err != nil { + return nil, "", "", err + } + + token, err := s.tokenProvider.GetAccessToken(ctx, account) + if err != nil { + return nil, "", "", infraerrors.Newf(http.StatusBadGateway, "GROK_QUOTA_TOKEN_UNAVAILABLE", "failed to acquire access token: %v", err) + } + if strings.TrimSpace(token) == "" { + return nil, "", "", infraerrors.New(http.StatusBadGateway, "GROK_QUOTA_TOKEN_UNAVAILABLE", "access token is empty") + } + + 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 "" + } + 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) { + if s == nil || s.accountRepo == nil { + return nil, infraerrors.New(http.StatusInternalServerError, "GROK_QUOTA_NOT_CONFIGURED", "grok quota service is not configured") + } + account, err := s.accountRepo.GetByID(ctx, accountID) + if err != nil { + return nil, infraerrors.Newf(http.StatusNotFound, "GROK_QUOTA_ACCOUNT_NOT_FOUND", "account not found: %v", err) + } + if account == nil { + return nil, infraerrors.New(http.StatusNotFound, "GROK_QUOTA_ACCOUNT_NOT_FOUND", "account not found") + } + if account.Platform != PlatformGrok { + return nil, infraerrors.New(http.StatusBadRequest, "GROK_QUOTA_INVALID_PLATFORM", "account is not a Grok account") + } + if account.Type != AccountTypeOAuth { + return nil, infraerrors.New(http.StatusBadRequest, "GROK_QUOTA_INVALID_TYPE", "account is not an OAuth account") + } + return account, nil +} + +func buildGrokQuotaProbeBody(account *Account) ([]byte, error) { + model := grokQuotaDefaultModel + if account != nil { + if mapped := strings.TrimSpace(account.GetMappedModel("grok")); mapped != "" { + model = mapped + } + } + return json.Marshal(map[string]any{ + "model": model, + "input": grokQuotaProbeInput, + "max_output_tokens": 1, + "store": false, + }) +} + +func maxInt(a, b int) int { + if a > b { + return a + } + return b +} diff --git a/backend/internal/service/grok_quota_service_test.go b/backend/internal/service/grok_quota_service_test.go new file mode 100644 index 0000000000..d1da2e50a5 --- /dev/null +++ b/backend/internal/service/grok_quota_service_test.go @@ -0,0 +1,302 @@ +//go:build unit + +package service + +import ( + "context" + "io" + "net/http" + "strings" + "testing" + "time" + + infraerrors "github.com/Wei-Shaw/sub2api/internal/pkg/errors" + "github.com/Wei-Shaw/sub2api/internal/pkg/xai" + "github.com/stretchr/testify/require" +) + +type grokQuotaAccountRepo struct { + *mockAccountRepoForPlatform + updates map[int64]map[string]any + tempUnschedCalls int + lastTempUnschedID int64 + lastTempUnschedUntil time.Time + lastTempUnschedReason string +} + +func (r *grokQuotaAccountRepo) UpdateExtra(_ context.Context, id int64, updates map[string]any) error { + if r.updates == nil { + r.updates = make(map[int64]map[string]any) + } + r.updates[id] = updates + return nil +} + +func (r *grokQuotaAccountRepo) SetTempUnschedulable(_ context.Context, id int64, until time.Time, reason string) error { + r.tempUnschedCalls++ + r.lastTempUnschedID = id + r.lastTempUnschedUntil = until + r.lastTempUnschedReason = reason + 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() + + account := &Account{ + ID: 42, + Platform: PlatformGrok, + Type: AccountTypeOAuth, + Concurrency: 1, + 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{42: account}, + }, + } + upstream := &httpUpstreamRecorder{resp: &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{ + "X-Ratelimit-Limit-Requests": []string{"10"}, + "X-Ratelimit-Remaining-Requests": []string{"7"}, + "X-Ratelimit-Reset-Requests": []string{"2000000000"}, + "X-Ratelimit-Limit-Tokens": []string{"1000"}, + "X-Ratelimit-Remaining-Tokens": []string{"900"}, + }, + Body: io.NopCloser(strings.NewReader(`{"id":"resp_probe"}`)), + }} + svc := NewGrokQuotaService(repo, nil, NewGrokTokenProvider(repo, nil), upstream) + + result, err := svc.ProbeUsage(context.Background(), 42) + require.NoError(t, err) + require.Equal(t, http.StatusOK, result.StatusCode) + require.True(t, result.HeadersObserved) + require.NotNil(t, result.Snapshot) + require.True(t, result.Snapshot.HeadersObserved) + require.Equal(t, "active_probe", result.Snapshot.ObservationSource) + require.NotEmpty(t, result.Snapshot.LastProbeAt) + require.NotEmpty(t, result.Snapshot.LastHeadersSeenAt) + require.NotNil(t, result.Snapshot.Requests) + require.EqualValues(t, 10, *result.Snapshot.Requests.Limit) + require.EqualValues(t, 7, *result.Snapshot.Requests.Remaining) + require.Equal(t, "https://api.x.ai/v1/responses", upstream.lastReq.URL.String()) + require.Equal(t, "Bearer access-token", upstream.lastReq.Header.Get("Authorization")) + require.Contains(t, string(upstream.lastBody), `"max_output_tokens":1`) + require.Contains(t, string(upstream.lastBody), `"store":false`) + 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() + + account := &Account{ + ID: 45, + Platform: PlatformGrok, + Type: AccountTypeOAuth, + Concurrency: 1, + 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{45: account}, + }, + } + upstream := &httpUpstreamRecorder{resp: &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{}, + Body: io.NopCloser(strings.NewReader(`{"id":"resp_probe"}`)), + }} + svc := NewGrokQuotaService(repo, nil, NewGrokTokenProvider(repo, nil), upstream) + + result, err := svc.ProbeUsage(context.Background(), 45) + require.NoError(t, err) + require.Equal(t, http.StatusOK, result.StatusCode) + require.False(t, result.HeadersObserved) + require.NotNil(t, result.Snapshot) + require.False(t, result.Snapshot.HeadersObserved) + require.Equal(t, "active_probe", result.Snapshot.ObservationSource) + require.NotEmpty(t, result.Snapshot.LastProbeAt) + require.Empty(t, result.Snapshot.LastHeadersSeenAt) + + stored, ok := repo.updates[45][grokQuotaSnapshotExtraKey].(*xai.QuotaSnapshot) + require.True(t, ok) + require.False(t, stored.HeadersObserved) + require.Equal(t, http.StatusOK, stored.StatusCode) +} + +func TestGrokQuotaServiceProbeUsageReturnsRateLimitedSnapshot(t *testing.T) { + t.Parallel() + + account := &Account{ + ID: 43, + Platform: PlatformGrok, + Type: AccountTypeOAuth, + 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{43: account}, + }, + } + upstream := &httpUpstreamRecorder{resp: &http.Response{ + StatusCode: http.StatusTooManyRequests, + Header: http.Header{"Retry-After": []string{"45"}}, + Body: io.NopCloser(strings.NewReader(`{"error":{"message":"rate limited"}}`)), + }} + svc := NewGrokQuotaService(repo, nil, NewGrokTokenProvider(repo, nil), upstream) + + result, err := svc.ProbeUsage(context.Background(), 43) + require.NoError(t, err) + require.Equal(t, http.StatusTooManyRequests, result.StatusCode) + require.NotNil(t, result.Snapshot) + require.NotNil(t, result.Snapshot.RetryAfterSeconds) + require.Equal(t, 45, *result.Snapshot.RetryAfterSeconds) +} + +func TestGrokQuotaServiceResetQuotaUnsupported(t *testing.T) { + t.Parallel() + + account := &Account{ + ID: 44, + Platform: PlatformGrok, + Type: AccountTypeOAuth, + } + repo := &grokQuotaAccountRepo{ + mockAccountRepoForPlatform: &mockAccountRepoForPlatform{ + accountsByID: map[int64]*Account{44: account}, + }, + } + svc := NewGrokQuotaService(repo, nil, nil, nil) + + _, err := svc.ResetQuota(context.Background(), 44) + require.Error(t, err) + require.Equal(t, http.StatusNotImplemented, infraerrors.Code(err)) + require.Equal(t, "GROK_QUOTA_RESET_UNSUPPORTED", infraerrors.Reason(err)) +} + +func TestShouldAutoPauseGrokAccountByQuota(t *testing.T) { + t.Parallel() + + zero := int64(0) + limit := int64(10) + resetFuture := time.Now().Add(time.Minute).Unix() + retryAfter := 30 + tests := []struct { + name string + snapshot xai.QuotaSnapshot + want bool + }{ + { + name: "remaining requests exhausted", + snapshot: xai.QuotaSnapshot{ + Requests: &xai.QuotaWindow{Limit: &limit, Remaining: &zero, ResetUnix: &resetFuture}, + UpdatedAt: time.Now().UTC().Format(time.RFC3339), + }, + want: true, + }, + { + name: "retry after active", + snapshot: xai.QuotaSnapshot{ + RetryAfterSeconds: &retryAfter, + UpdatedAt: time.Now().UTC().Format(time.RFC3339), + }, + want: true, + }, + { + name: "retry after expired", + snapshot: xai.QuotaSnapshot{ + RetryAfterSeconds: &retryAfter, + UpdatedAt: time.Now().Add(-time.Duration(retryAfter+1) * time.Second).UTC().Format(time.RFC3339), + }, + want: false, + }, + { + name: "stale snapshot ignored", + snapshot: xai.QuotaSnapshot{ + Requests: &xai.QuotaWindow{Limit: &limit, Remaining: &zero, ResetUnix: &resetFuture}, + UpdatedAt: time.Now().Add(-3 * time.Hour).UTC().Format(time.RFC3339), + }, + want: false, + }, + } + + for _, tt := range tests { + tt := tt + t.Run(tt.name, func(t *testing.T) { + t.Parallel() + account := &Account{ + Platform: PlatformGrok, + Type: AccountTypeOAuth, + Extra: map[string]any{ + grokQuotaSnapshotExtraKey: tt.snapshot, + }, + } + got, _ := shouldAutoPauseGrokAccountByQuota(account) + require.Equal(t, tt.want, got) + }) + } +} diff --git a/backend/internal/service/grok_token_provider.go b/backend/internal/service/grok_token_provider.go new file mode 100644 index 0000000000..f19ee42b88 --- /dev/null +++ b/backend/internal/service/grok_token_provider.go @@ -0,0 +1,165 @@ +package service + +import ( + "context" + "errors" + "log/slog" + "strconv" + "strings" + "time" + + "github.com/Wei-Shaw/sub2api/internal/util/logredact" +) + +const ( + grokTokenCacheSkew = 5 * time.Minute + grokRequestRefreshTimeout = 8 * time.Second + grokTokenProviderLogComponent = "grok_token_provider" + grokTempUnschedulableErrorCode = "token_refresh_failed" +) + +type GrokTokenCache = GeminiTokenCache + +type GrokTokenProvider struct { + accountRepo AccountRepository + tokenCache GrokTokenCache + refreshAPI *OAuthRefreshAPI + executor OAuthRefreshExecutor + refreshPolicy ProviderRefreshPolicy + tempUnschedCache TempUnschedCache +} + +func NewGrokTokenProvider( + accountRepo AccountRepository, + tokenCache GrokTokenCache, +) *GrokTokenProvider { + return &GrokTokenProvider{ + accountRepo: accountRepo, + tokenCache: tokenCache, + refreshPolicy: AntigravityProviderRefreshPolicy(), + } +} + +func (p *GrokTokenProvider) SetRefreshAPI(api *OAuthRefreshAPI, executor OAuthRefreshExecutor) { + p.refreshAPI = api + p.executor = executor +} + +func (p *GrokTokenProvider) SetRefreshPolicy(policy ProviderRefreshPolicy) { + p.refreshPolicy = policy +} + +func (p *GrokTokenProvider) SetTempUnschedCache(cache TempUnschedCache) { + p.tempUnschedCache = cache +} + +func (p *GrokTokenProvider) GetAccessToken(ctx context.Context, account *Account) (string, error) { + if account == nil { + return "", errors.New("account is nil") + } + if account.Platform != PlatformGrok || account.Type != AccountTypeOAuth { + return "", errors.New("not a grok oauth account") + } + + cacheKey := GrokTokenCacheKey(account) + if p.tokenCache != nil { + if token, err := p.tokenCache.GetAccessToken(ctx, cacheKey); err == nil && strings.TrimSpace(token) != "" { + return token, nil + } + } + + expiresAt := account.GetCredentialAsTime("expires_at") + needsRefresh := expiresAt == nil || time.Until(*expiresAt) <= grokTokenRefreshSkew + if needsRefresh && strings.TrimSpace(account.GetGrokRefreshToken()) == "" { + if expiresAt == nil || !time.Now().Before(*expiresAt) { + return "", errors.New("grok access_token expired and refresh_token is missing") + } + needsRefresh = false + } + if needsRefresh && p.refreshAPI != nil && p.executor != nil { + refreshCtx, cancel := context.WithTimeout(ctx, grokRequestRefreshTimeout) + defer cancel() + result, err := p.refreshAPI.RefreshIfNeeded(refreshCtx, account, p.executor, grokTokenRefreshSkew) + if err != nil { + p.markTempUnschedulable(account, err) + if p.refreshPolicy.OnRefreshError == ProviderRefreshErrorReturn { + return "", err + } + } else if !result.LockHeld && result.Account != nil { + account = result.Account + expiresAt = account.GetCredentialAsTime("expires_at") + } + } + + accessToken := account.GetGrokAccessToken() + if strings.TrimSpace(accessToken) == "" { + return "", errors.New("access_token not found in credentials") + } + + if p.tokenCache != nil { + latestAccount, isStale := CheckTokenVersion(ctx, account, p.accountRepo) + if isStale && latestAccount != nil { + accessToken = latestAccount.GetGrokAccessToken() + if strings.TrimSpace(accessToken) == "" { + return "", errors.New("access_token not found after version check") + } + } else { + ttl := 30 * time.Minute + if expiresAt != nil { + until := time.Until(*expiresAt) + switch { + case until > grokTokenCacheSkew: + ttl = until - grokTokenCacheSkew + case until > 0: + ttl = until + default: + ttl = time.Minute + } + } + _ = p.tokenCache.SetAccessToken(ctx, cacheKey, accessToken, ttl) + } + } + + return accessToken, nil +} + +func (p *GrokTokenProvider) markTempUnschedulable(account *Account, refreshErr error) { + if p == nil || p.accountRepo == nil || account == nil { + return + } + now := time.Now() + until := now.Add(tokenRefreshTempUnschedDuration) + 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) + return + } + if p.tempUnschedCache != nil { + state := &TempUnschedState{ + UntilUnix: until.Unix(), + TriggeredAtUnix: now.Unix(), + ErrorMessage: grokTempUnschedulableErrorCode + ": " + reason, + } + if err := p.tempUnschedCache.SetTempUnsched(bgCtx, account.ID, state); err != nil { + slog.Warn(grokTokenProviderLogComponent+".temp_unsched_cache_set_failed", "account_id", account.ID, "error", err) + } + } +} + +func GrokTokenCacheKey(account *Account) string { + if account == nil { + return "grok:account:0" + } + return "grok:account:" + strconv.FormatInt(account.ID, 10) +} diff --git a/backend/internal/service/grok_token_provider_test.go b/backend/internal/service/grok_token_provider_test.go new file mode 100644 index 0000000000..1f7e2d8e0e --- /dev/null +++ b/backend/internal/service/grok_token_provider_test.go @@ -0,0 +1,131 @@ +//go:build unit + +package service + +import ( + "context" + "errors" + "testing" + "time" + + "github.com/Wei-Shaw/sub2api/internal/pkg/xai" + "github.com/stretchr/testify/require" +) + +type grokTokenCacheForProviderTest struct { + token string + setKey string + setToken string + setTTL time.Duration + lockResult bool + releaseCalls int +} + +func (c *grokTokenCacheForProviderTest) GetAccessToken(context.Context, string) (string, error) { + if c.token == "" { + return "", errors.New("not cached") + } + return c.token, nil +} + +func (c *grokTokenCacheForProviderTest) SetAccessToken(_ context.Context, key string, token string, ttl time.Duration) error { + c.setKey = key + c.setToken = token + c.setTTL = ttl + return nil +} + +func (c *grokTokenCacheForProviderTest) DeleteAccessToken(context.Context, string) error { + return nil +} + +func (c *grokTokenCacheForProviderTest) AcquireRefreshLock(context.Context, string, time.Duration) (bool, error) { + return c.lockResult, nil +} + +func (c *grokTokenCacheForProviderTest) ReleaseRefreshLock(context.Context, string) error { + c.releaseCalls++ + return nil +} + +func TestGrokTokenProviderRefreshesExpiredTokenOnRequestPath(t *testing.T) { + t.Setenv(xai.EnvBaseURL, xai.DefaultCLIBaseURL) + + expiredAt := time.Now().Add(-time.Minute).UTC().Format(time.RFC3339) + account := &Account{ + ID: 54, + Platform: PlatformGrok, + Type: AccountTypeOAuth, + Credentials: map[string]any{ + "access_token": "expired-access-token", + "refresh_token": "refresh-token", + "expires_at": expiredAt, + "base_url": xai.DefaultCLIBaseURL, + "client_id": "client-id", + }, + } + repo := &tokenRefreshAccountRepo{} + repo.accountsByID = map[int64]*Account{54: account} + cache := &grokTokenCacheForProviderTest{lockResult: true} + oauthSvc := NewGrokOAuthService(nil, &grokOAuthClientStub{ + refreshResponse: &xai.TokenResponse{ + AccessToken: "new-access-token", + TokenType: "Bearer", + ExpiresIn: 3600, + }, + }) + defer oauthSvc.Stop() + + provider := NewGrokTokenProvider(repo, cache) + provider.SetRefreshAPI(NewOAuthRefreshAPI(repo, cache), NewGrokTokenRefresher(oauthSvc)) + + token, err := provider.GetAccessToken(context.Background(), account) + require.NoError(t, err) + require.Equal(t, "new-access-token", token) + require.Equal(t, 1, repo.updateCredentialsCalls) + require.Equal(t, "new-access-token", repo.accountsByID[54].GetGrokAccessToken()) + require.Equal(t, "refresh-token", repo.accountsByID[54].GetGrokRefreshToken()) + require.Equal(t, xai.DefaultCLIBaseURL, repo.accountsByID[54].GetGrokBaseURL()) + require.Equal(t, "grok:account:54", cache.setKey) + require.Equal(t, "new-access-token", cache.setToken) + require.Greater(t, cache.setTTL, time.Duration(0)) + require.Equal(t, 1, cache.releaseCalls) +} + +func TestGrokTokenProviderRefreshFailureUnschedulesWithRedactedReason(t *testing.T) { + expiredAt := time.Now().Add(-time.Minute).UTC().Format(time.RFC3339) + account := &Account{ + ID: 55, + Platform: PlatformGrok, + Type: AccountTypeOAuth, + Credentials: map[string]any{ + "access_token": "expired-access-token", + "refresh_token": "refresh-token", + "expires_at": expiredAt, + "base_url": xai.DefaultCLIBaseURL, + }, + } + repo := &tokenRefreshAccountRepo{} + repo.accountsByID = map[int64]*Account{55: account} + cache := &grokTokenCacheForProviderTest{lockResult: true} + tempCache := &tempUnschedCacheStub{} + provider := NewGrokTokenProvider(repo, cache) + provider.SetRefreshAPI(NewOAuthRefreshAPI(repo, cache), &tokenRefresherStub{ + err: errors.New("temporary refresh failure access_token=leaked-access refresh_token=leaked-refresh"), + }) + provider.SetTempUnschedCache(tempCache) + + token, err := provider.GetAccessToken(context.Background(), account) + require.Error(t, err) + require.Empty(t, token) + require.Equal(t, 1, repo.setTempUnschedCalls) + require.Equal(t, 0, repo.setErrorCalls) + require.Contains(t, repo.lastTempUnschedReason, "access_token=***") + require.Contains(t, repo.lastTempUnschedReason, "refresh_token=***") + require.NotContains(t, repo.lastTempUnschedReason, "leaked-access") + require.NotContains(t, repo.lastTempUnschedReason, "leaked-refresh") + require.Equal(t, 1, tempCache.setCalls) + require.NotNil(t, tempCache.lastState) + require.NotContains(t, tempCache.lastState.ErrorMessage, "leaked-access") + require.NotContains(t, tempCache.lastState.ErrorMessage, "leaked-refresh") +} diff --git a/backend/internal/service/grok_token_refresher.go b/backend/internal/service/grok_token_refresher.go new file mode 100644 index 0000000000..92d88cc058 --- /dev/null +++ b/backend/internal/service/grok_token_refresher.go @@ -0,0 +1,56 @@ +package service + +import ( + "context" + "errors" + "strings" + "time" +) + +const grokTokenRefreshSkew = time.Hour + +type GrokTokenRefresher struct { + grokOAuthService GrokOAuthTokenService +} + +func NewGrokTokenRefresher(grokOAuthService GrokOAuthTokenService) *GrokTokenRefresher { + return &GrokTokenRefresher{grokOAuthService: grokOAuthService} +} + +func (r *GrokTokenRefresher) CacheKey(account *Account) string { + return GrokTokenCacheKey(account) +} + +func (r *GrokTokenRefresher) CanRefresh(account *Account) bool { + return account != nil && account.Platform == PlatformGrok && account.Type == AccountTypeOAuth +} + +func (r *GrokTokenRefresher) NeedsRefresh(account *Account, refreshWindow time.Duration) bool { + if account == nil || strings.TrimSpace(account.GetGrokRefreshToken()) == "" { + return false + } + expiresAt := account.GetCredentialAsTime("expires_at") + if expiresAt == nil { + return true + } + if refreshWindow < grokTokenRefreshSkew { + refreshWindow = grokTokenRefreshSkew + } + return time.Until(*expiresAt) < refreshWindow +} + +func (r *GrokTokenRefresher) Refresh(ctx context.Context, account *Account) (map[string]any, error) { + if r == nil || r.grokOAuthService == nil { + return nil, errors.New("grok oauth service is not configured") + } + tokenInfo, err := r.grokOAuthService.RefreshAccountToken(ctx, account) + if err != nil { + return nil, err + } + newCredentials := r.grokOAuthService.BuildAccountCredentials(tokenInfo) + newCredentials = MergeCredentials(account.Credentials, newCredentials) + if baseURL := strings.TrimSpace(account.GetCredential("base_url")); baseURL != "" { + newCredentials["base_url"] = baseURL + } + return newCredentials, nil +} diff --git a/backend/internal/service/oauth_service.go b/backend/internal/service/oauth_service.go index 0931f9ce81..1369dd9e89 100644 --- a/backend/internal/service/oauth_service.go +++ b/backend/internal/service/oauth_service.go @@ -8,6 +8,7 @@ import ( "github.com/Wei-Shaw/sub2api/internal/pkg/oauth" "github.com/Wei-Shaw/sub2api/internal/pkg/openai" + "github.com/Wei-Shaw/sub2api/internal/pkg/xai" ) // OpenAIOAuthClient interface for OpenAI OAuth operations @@ -17,6 +18,18 @@ type OpenAIOAuthClient interface { RefreshTokenWithClientID(ctx context.Context, refreshToken, proxyURL string, clientID string) (*openai.TokenResponse, error) } +// GrokOAuthClient interface for xAI/Grok OAuth operations. +type GrokOAuthClient interface { + ExchangeCode(ctx context.Context, code, codeVerifier, redirectURI, proxyURL, clientID string) (*xai.TokenResponse, error) + 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_account_runtime_block_fastpath.go b/backend/internal/service/openai_account_runtime_block_fastpath.go index bdd97f479a..0a17f3b938 100644 --- a/backend/internal/service/openai_account_runtime_block_fastpath.go +++ b/backend/internal/service/openai_account_runtime_block_fastpath.go @@ -27,8 +27,12 @@ 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 + return account != nil && (account.Platform == PlatformOpenAI || account.Platform == PlatformGrok) } func (s *OpenAIGatewayService) handleOpenAIAccountUpstreamError(ctx context.Context, account *Account, statusCode int, headers http.Header, responseBody []byte, requestedModel ...string) bool { @@ -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 } diff --git a/backend/internal/service/openai_account_runtime_block_fastpath_test.go b/backend/internal/service/openai_account_runtime_block_fastpath_test.go index 3784dd3386..ff5d604fe4 100644 --- a/backend/internal/service/openai_account_runtime_block_fastpath_test.go +++ b/backend/internal/service/openai_account_runtime_block_fastpath_test.go @@ -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)) +} diff --git a/backend/internal/service/openai_account_scheduler.go b/backend/internal/service/openai_account_scheduler.go index 59da55d3fe..2277b1d904 100644 --- a/backend/internal/service/openai_account_scheduler.go +++ b/backend/internal/service/openai_account_scheduler.go @@ -39,6 +39,7 @@ var openAIAdvancedSchedulerSettingSF singleflight.Group type OpenAIAccountScheduleRequest struct { GroupID *int64 + Platform string SessionHash string StickyAccountID int64 PreserveStickyBinding bool @@ -270,7 +271,7 @@ func (s *defaultOpenAIAccountScheduler) Select( }() previousResponseID := strings.TrimSpace(req.PreviousResponseID) - if previousResponseID != "" { + if previousResponseID != "" && normalizeOpenAICompatiblePlatform(req.Platform) == PlatformOpenAI { selection, err := s.service.selectAccountByPreviousResponseIDForCapability( ctx, req.GroupID, @@ -364,7 +365,7 @@ func (s *defaultOpenAIAccountScheduler) selectBySessionHash( _ = s.service.deleteStickySessionAccountID(ctx, req.GroupID, sessionHash) return nil, false, nil } - if shouldClearStickySession(account, req.RequestedModel) || !account.IsOpenAI() || !account.IsSchedulable() { + if shouldClearStickySession(account, req.RequestedModel) || account.Platform != normalizeOpenAICompatiblePlatform(req.Platform) || !account.IsOpenAICompatible() || !account.IsSchedulable() { _ = s.service.deleteStickySessionAccountID(ctx, req.GroupID, sessionHash) return nil, false, nil } @@ -375,7 +376,7 @@ func (s *defaultOpenAIAccountScheduler) selectBySessionHash( _ = s.service.deleteStickySessionAccountID(ctx, req.GroupID, sessionHash) return nil, false, nil } - account = s.service.recheckSelectedOpenAIAccountFromDB(ctx, account, req.RequestedModel, req.RequireCompact, req.RequiredCapability) + account = s.service.recheckSelectedOpenAIAccountFromDB(ctx, account, req.Platform, req.RequestedModel, req.RequireCompact, req.RequiredCapability) if account == nil || !openAIStickyAccountMatchesGroup(account, req.GroupID) || !s.isAccountTransportCompatible(account, req.RequiredTransport) { _ = s.service.deleteStickySessionAccountID(ctx, req.GroupID, sessionHash) return nil, false, nil @@ -897,11 +898,11 @@ func (s *defaultOpenAIAccountScheduler) tryAcquireOpenAISelectionOrder( compactBlocked := false for i := 0; i < len(selectionOrder); i++ { candidate := selectionOrder[i] - fresh := s.service.resolveFreshSchedulableOpenAIAccount(ctx, candidate.account, req.RequestedModel, false, req.RequiredCapability) + fresh := s.service.resolveFreshSchedulableOpenAIAccount(ctx, candidate.account, req.Platform, req.RequestedModel, false, req.RequiredCapability) if fresh == nil || !s.isAccountTransportCompatible(fresh, req.RequiredTransport) || !s.isAccountRequestCompatible(ctx, fresh, req) { continue } - fresh = s.service.recheckSelectedOpenAIAccountFromDB(ctx, fresh, req.RequestedModel, false, req.RequiredCapability) + fresh = s.service.recheckSelectedOpenAIAccountFromDB(ctx, fresh, req.Platform, req.RequestedModel, false, req.RequiredCapability) if fresh == nil || !s.isAccountTransportCompatible(fresh, req.RequiredTransport) || !s.isAccountRequestCompatible(ctx, fresh, req) { continue } @@ -931,7 +932,7 @@ func (s *defaultOpenAIAccountScheduler) selectByLoadBalance( ctx context.Context, req OpenAIAccountScheduleRequest, ) (*AccountSelectionResult, int, int, float64, error) { - accounts, err := s.service.listSchedulableAccounts(ctx, req.GroupID) + accounts, err := s.service.listSchedulableAccounts(ctx, req.GroupID, req.Platform) if err != nil { return nil, 0, 0, 0, err } @@ -954,7 +955,7 @@ func (s *defaultOpenAIAccountScheduler) selectByLoadBalance( continue } } - if !account.IsSchedulable() || !account.IsOpenAI() { + if !account.IsSchedulable() || account.Platform != normalizeOpenAICompatiblePlatform(req.Platform) || !account.IsOpenAICompatible() { continue } if s.service.isOpenAIAccountRuntimeBlocked(account) { @@ -1036,11 +1037,11 @@ func (s *defaultOpenAIAccountScheduler) selectByLoadBalance( cfg := s.service.schedulingConfig() // WaitPlan.MaxConcurrency 使用 Concurrency(非 EffectiveLoadFactor),因为 WaitPlan 控制的是 Redis 实际并发槽位等待。 for _, candidate := range selectionOrder { - fresh := s.service.resolveFreshSchedulableOpenAIAccount(ctx, candidate.account, req.RequestedModel, false, req.RequiredCapability) + fresh := s.service.resolveFreshSchedulableOpenAIAccount(ctx, candidate.account, req.Platform, req.RequestedModel, false, req.RequiredCapability) if fresh == nil || !s.isAccountTransportCompatible(fresh, req.RequiredTransport) || !s.isAccountRequestCompatible(ctx, fresh, req) { continue } - fresh = s.service.recheckSelectedOpenAIAccountFromDB(ctx, fresh, req.RequestedModel, false, req.RequiredCapability) + fresh = s.service.recheckSelectedOpenAIAccountFromDB(ctx, fresh, req.Platform, req.RequestedModel, false, req.RequiredCapability) if fresh == nil || !s.isAccountTransportCompatible(fresh, req.RequiredTransport) || !s.isAccountRequestCompatible(ctx, fresh, req) { continue } @@ -1217,7 +1218,7 @@ func (s *OpenAIGatewayService) SelectAccountWithScheduler( requiredTransport OpenAIUpstreamTransport, requireCompact bool, ) (*AccountSelectionResult, OpenAIAccountScheduleDecision, error) { - return s.selectAccountWithScheduler(ctx, groupID, previousResponseID, sessionHash, requestedModel, excludedIDs, requiredTransport, "", "", requireCompact) + return s.selectAccountWithScheduler(ctx, groupID, previousResponseID, sessionHash, requestedModel, excludedIDs, requiredTransport, "", "", requireCompact, PlatformOpenAI) } func (s *OpenAIGatewayService) SelectAccountWithSchedulerForCapability( @@ -1230,8 +1231,13 @@ func (s *OpenAIGatewayService) SelectAccountWithSchedulerForCapability( requiredTransport OpenAIUpstreamTransport, requiredCapability OpenAIEndpointCapability, requireCompact bool, + platformOverride ...string, ) (*AccountSelectionResult, OpenAIAccountScheduleDecision, error) { - return s.selectAccountWithScheduler(ctx, groupID, previousResponseID, sessionHash, requestedModel, excludedIDs, requiredTransport, requiredCapability, "", requireCompact) + platform := PlatformOpenAI + if len(platformOverride) > 0 { + platform = platformOverride[0] + } + return s.selectAccountWithScheduler(ctx, groupID, previousResponseID, sessionHash, requestedModel, excludedIDs, requiredTransport, requiredCapability, "", requireCompact, platform) } func (s *OpenAIGatewayService) SelectAccountWithSchedulerForImages( @@ -1242,13 +1248,13 @@ func (s *OpenAIGatewayService) SelectAccountWithSchedulerForImages( excludedIDs map[int64]struct{}, requiredCapability OpenAIImagesCapability, ) (*AccountSelectionResult, OpenAIAccountScheduleDecision, error) { - selection, decision, err := s.selectAccountWithScheduler(ctx, groupID, "", sessionHash, requestedModel, excludedIDs, OpenAIUpstreamTransportHTTPSSE, "", requiredCapability, false) + selection, decision, err := s.selectAccountWithScheduler(ctx, groupID, "", sessionHash, requestedModel, excludedIDs, OpenAIUpstreamTransportHTTPSSE, "", requiredCapability, false, PlatformOpenAI) if err == nil && selection != nil && selection.Account != nil { return selection, decision, nil } // 如果要求 native 能力(如指定了模型)但没有可用的 APIKey 账号,回退到 basic(OAuth 账号) if requiredCapability == OpenAIImagesCapabilityNative { - return s.selectAccountWithScheduler(ctx, groupID, "", sessionHash, requestedModel, excludedIDs, OpenAIUpstreamTransportHTTPSSE, "", OpenAIImagesCapabilityBasic, false) + return s.selectAccountWithScheduler(ctx, groupID, "", sessionHash, requestedModel, excludedIDs, OpenAIUpstreamTransportHTTPSSE, "", OpenAIImagesCapabilityBasic, false, PlatformOpenAI) } return selection, decision, err } @@ -1264,8 +1270,10 @@ func (s *OpenAIGatewayService) selectAccountWithScheduler( requiredCapability OpenAIEndpointCapability, requiredImageCapability OpenAIImagesCapability, requireCompact bool, + platform string, ) (*AccountSelectionResult, OpenAIAccountScheduleDecision, error) { ctx = s.withOpenAIQuotaAutoPauseContext(ctx) + platform = normalizeOpenAICompatiblePlatform(platform) decision := OpenAIAccountScheduleDecision{} scheduler := s.getOpenAIAccountScheduler(ctx) if scheduler == nil { @@ -1273,7 +1281,7 @@ func (s *OpenAIGatewayService) selectAccountWithScheduler( if requiredTransport == OpenAIUpstreamTransportAny || requiredTransport == OpenAIUpstreamTransportHTTPSSE { effectiveExcludedIDs := cloneExcludedAccountIDs(excludedIDs) for { - selection, err := s.selectAccountWithLoadAwareness(ctx, groupID, sessionHash, requestedModel, effectiveExcludedIDs, requireCompact, requiredCapability) + selection, err := s.selectAccountWithLoadAwareness(ctx, groupID, platform, sessionHash, requestedModel, effectiveExcludedIDs, requireCompact, requiredCapability) if err != nil { return nil, decision, err } @@ -1298,7 +1306,7 @@ func (s *OpenAIGatewayService) selectAccountWithScheduler( effectiveExcludedIDs := cloneExcludedAccountIDs(excludedIDs) for { - selection, err := s.selectAccountWithLoadAwareness(ctx, groupID, sessionHash, requestedModel, effectiveExcludedIDs, requireCompact, requiredCapability) + selection, err := s.selectAccountWithLoadAwareness(ctx, groupID, platform, sessionHash, requestedModel, effectiveExcludedIDs, requireCompact, requiredCapability) if err != nil { return nil, decision, err } @@ -1338,6 +1346,7 @@ func (s *OpenAIGatewayService) selectAccountWithScheduler( return scheduler.Select(ctx, OpenAIAccountScheduleRequest{ GroupID: groupID, + Platform: platform, SessionHash: sessionHash, StickyAccountID: stickyAccountID, PreviousResponseID: previousResponseID, diff --git a/backend/internal/service/openai_account_scheduler_test.go b/backend/internal/service/openai_account_scheduler_test.go index 4d2fbd7d70..6d8e38d0bc 100644 --- a/backend/internal/service/openai_account_scheduler_test.go +++ b/backend/internal/service/openai_account_scheduler_test.go @@ -475,6 +475,50 @@ func TestOpenAIGatewayService_SelectAccountWithScheduler_DefaultDisabled_Embeddi require.Equal(t, openAIAccountScheduleLayerLoadBalance, decision.Layer) } +func TestOpenAIGatewayService_SelectAccountWithScheduler_DefaultDisabled_AllowsGrokChatAccount(t *testing.T) { + resetOpenAIAdvancedSchedulerSettingCacheForTest() + + ctx := context.Background() + groupID := int64(10113) + accounts := []Account{ + { + ID: 36041, + Platform: PlatformGrok, + Type: AccountTypeOAuth, + Status: StatusActive, + Schedulable: true, + Concurrency: 1, + Priority: 0, + }, + } + cfg := &config.Config{} + cfg.Gateway.Scheduling.LoadBatchEnabled = false + svc := &OpenAIGatewayService{ + accountRepo: schedulerTestOpenAIAccountRepo{accounts: accounts}, + cache: &schedulerTestGatewayCache{}, + cfg: cfg, + concurrencyService: NewConcurrencyService(schedulerTestConcurrencyCache{}), + } + + selection, decision, err := svc.SelectAccountWithSchedulerForCapability( + ctx, + &groupID, + "", + "", + "grok-4.3", + nil, + OpenAIUpstreamTransportAny, + OpenAIEndpointCapabilityChatCompletions, + false, + PlatformGrok, + ) + require.NoError(t, err) + require.NotNil(t, selection) + require.NotNil(t, selection.Account) + require.Equal(t, int64(36041), selection.Account.ID) + require.Equal(t, openAIAccountScheduleLayerLoadBalance, decision.Layer) +} + func TestOpenAIGatewayService_SelectAccountWithScheduler_EnabledUsesAdvancedPreviousResponseRouting(t *testing.T) { resetOpenAIAdvancedSchedulerSettingCacheForTest() diff --git a/backend/internal/service/openai_gateway_chat_completions.go b/backend/internal/service/openai_gateway_chat_completions.go index 50135dd0c5..6035dc4ccd 100644 --- a/backend/internal/service/openai_gateway_chat_completions.go +++ b/backend/internal/service/openai_gateway_chat_completions.go @@ -74,6 +74,10 @@ func (s *OpenAIGatewayService) ForwardAsChatCompletions( return nil, errors.New("codex_cli_only restriction: only codex official clients are allowed") } + if account.Platform == PlatformGrok { + return s.forwardAsRawChatCompletions(ctx, c, account, body, defaultMappedModel) + } + // 入口分流:APIKey 账号 + 强制或已探测确认上游不支持 Responses,走 CC 直转。 // 自动模式下标记缺失(未探测)按"现状即证据"原则继续走下方原 Responses 转换路径。 if account.Type == AccountTypeAPIKey && !openai_compat.ShouldUseResponsesAPI(account.Extra) { diff --git a/backend/internal/service/openai_gateway_chat_completions_raw.go b/backend/internal/service/openai_gateway_chat_completions_raw.go index eef980128b..6bcb6718b7 100644 --- a/backend/internal/service/openai_gateway_chat_completions_raw.go +++ b/backend/internal/service/openai_gateway_chat_completions_raw.go @@ -14,6 +14,7 @@ import ( "github.com/Wei-Shaw/sub2api/internal/pkg/apicompat" "github.com/Wei-Shaw/sub2api/internal/pkg/logger" + "github.com/Wei-Shaw/sub2api/internal/pkg/xai" "github.com/Wei-Shaw/sub2api/internal/util/responseheaders" "github.com/gin-gonic/gin" "github.com/tidwall/gjson" @@ -121,19 +122,18 @@ func (s *OpenAIGatewayService) forwardAsRawChatCompletions( ) // 5. Build upstream request - apiKey := account.GetOpenAIApiKey() - if apiKey == "" { - return nil, fmt.Errorf("account %d missing api_key", account.ID) - } - baseURL := account.GetOpenAIBaseURL() - if baseURL == "" { - baseURL = "https://api.openai.com" - } - validatedURL, err := s.validateUpstreamBaseURL(baseURL) + token, tokenKind, err := s.GetAccessToken(ctx, account) if err != nil { - return nil, fmt.Errorf("invalid base_url: %w", err) + return nil, err + } + if strings.TrimSpace(token) == "" { + return nil, fmt.Errorf("account %d missing %s credential", account.ID, tokenKind) + } + + targetURL, err := s.rawChatCompletionsURL(account) + if err != nil { + return nil, err } - targetURL := buildOpenAIChatCompletionsURL(validatedURL) upstreamCtx, releaseUpstreamCtx := detachUpstreamContext(ctx) upstreamReq, err := http.NewRequestWithContext(upstreamCtx, http.MethodPost, targetURL, bytes.NewReader(upstreamBody)) @@ -143,7 +143,7 @@ func (s *OpenAIGatewayService) forwardAsRawChatCompletions( } upstreamReq = upstreamReq.WithContext(WithHTTPUpstreamProfile(upstreamReq.Context(), HTTPUpstreamProfileOpenAI)) upstreamReq.Header.Set("Content-Type", "application/json") - upstreamReq.Header.Set("Authorization", "Bearer "+apiKey) + upstreamReq.Header.Set("Authorization", "Bearer "+token) if clientStream { upstreamReq.Header.Set("Accept", "text/event-stream") } else { @@ -162,6 +162,8 @@ func (s *OpenAIGatewayService) forwardAsRawChatCompletions( customUA := account.GetOpenAIUserAgent() if customUA != "" { upstreamReq.Header.Set("user-agent", customUA) + } else if account.Platform == PlatformGrok { + upstreamReq.Header.Set("user-agent", "sub2api-grok/1.0") } // 6. Send request @@ -180,9 +182,32 @@ func (s *OpenAIGatewayService) forwardAsRawChatCompletions( respBody := s.readUpstreamErrorBody(resp) _ = resp.Body.Close() resp.Body = io.NopCloser(bytes.NewReader(respBody)) + if account.Platform == PlatformGrok { + s.updateGrokUsageSnapshot(ctx, account.ID, xai.ParseQuotaHeaders(resp.Header, resp.StatusCode)) + } upstreamMsg := strings.TrimSpace(extractUpstreamErrorMessage(respBody)) upstreamMsg = sanitizeUpstreamErrorMessage(upstreamMsg) + if account.Platform == PlatformGrok { + appendOpsUpstreamError(c, OpsUpstreamErrorEvent{ + Platform: account.Platform, + AccountID: account.ID, + AccountName: account.Name, + UpstreamStatusCode: resp.StatusCode, + UpstreamRequestID: firstNonEmpty(resp.Header.Get("x-request-id"), resp.Header.Get("xai-request-id")), + Kind: "failover", + Message: upstreamMsg, + }) + s.handleGrokAccountUpstreamError(ctx, account, resp.StatusCode, resp.Header, respBody) + if s.shouldFailoverUpstreamError(resp.StatusCode) { + return nil, &UpstreamFailoverError{ + StatusCode: resp.StatusCode, + ResponseBody: respBody, + RetryableOnSameAccount: account.IsPoolMode() && account.IsPoolModeRetryableStatus(resp.StatusCode), + } + } + return s.handleChatCompletionsErrorResponse(resp, c, account, billingModel) + } if s.shouldFailoverOpenAIUpstreamResponse(resp.StatusCode, upstreamMsg, respBody) { upstreamDetail := "" if s.cfg != nil && s.cfg.Gateway.LogUpstreamErrorBody { @@ -212,6 +237,10 @@ func (s *OpenAIGatewayService) forwardAsRawChatCompletions( return s.handleChatCompletionsErrorResponse(resp, c, account, billingModel) } + if account.Platform == PlatformGrok { + s.updateGrokUsageSnapshot(ctx, account.ID, xai.ParseQuotaHeaders(resp.Header, resp.StatusCode)) + } + // 8. Forward response if clientStream { return s.streamRawChatCompletions(c, resp, account, originalModel, billingModel, upstreamModel, reasoningEffort, serviceTier, startTime, len(body)) @@ -219,6 +248,26 @@ func (s *OpenAIGatewayService) forwardAsRawChatCompletions( return s.bufferRawChatCompletions(c, resp, originalModel, billingModel, upstreamModel, reasoningEffort, serviceTier, startTime) } +func (s *OpenAIGatewayService) rawChatCompletionsURL(account *Account) (string, error) { + if account.Platform == PlatformGrok { + targetURL, err := xai.BuildChatCompletionsURL(account.GetGrokBaseURL()) + if err != nil { + return "", fmt.Errorf("invalid grok base_url: %w", err) + } + return targetURL, nil + } + + baseURL := account.GetOpenAIBaseURL() + if baseURL == "" { + baseURL = "https://api.openai.com" + } + validatedURL, err := s.validateUpstreamBaseURL(baseURL) + if err != nil { + return "", fmt.Errorf("invalid base_url: %w", err) + } + return buildOpenAIChatCompletionsURL(validatedURL), nil +} + // streamRawChatCompletions 透传上游 CC SSE 流到客户端,并提取 usage(包括 // 末尾 [DONE] 之前的 chunk 中的 usage 字段,按 OpenAI CC 协议)。 // diff --git a/backend/internal/service/openai_gateway_grok.go b/backend/internal/service/openai_gateway_grok.go new file mode 100644 index 0000000000..eb3fbb66e6 --- /dev/null +++ b/backend/internal/service/openai_gateway_grok.go @@ -0,0 +1,230 @@ +package service + +import ( + "bytes" + "context" + "encoding/json" + "fmt" + "io" + "net/http" + "strings" + "time" + + "github.com/Wei-Shaw/sub2api/internal/pkg/xai" + "github.com/gin-gonic/gin" + "github.com/tidwall/gjson" + "github.com/tidwall/sjson" +) + +func (s *OpenAIGatewayService) forwardGrokResponses( + ctx context.Context, + c *gin.Context, + account *Account, + body []byte, + originalModel string, + reqStream bool, + startTime time.Time, +) (*OpenAIForwardResult, error) { + if account.Type != AccountTypeOAuth { + return nil, fmt.Errorf("grok account type %s is not supported by subscription forwarding", account.Type) + } + + upstreamModel := account.GetMappedModel(originalModel) + if strings.TrimSpace(upstreamModel) == "" { + upstreamModel = "grok-4.3" + } + patchedBody, err := patchGrokResponsesBody(body, upstreamModel) + if err != nil { + return nil, err + } + + token, _, err := s.GetAccessToken(ctx, account) + if err != nil { + return nil, err + } + + upstreamCtx, releaseUpstreamCtx := detachUpstreamContext(ctx) + defer releaseUpstreamCtx() + upstreamReq, err := buildGrokResponsesRequest(upstreamCtx, c, account, patchedBody, token) + if err != nil { + return nil, err + } + + proxyURL := "" + if account.ProxyID != nil && account.Proxy != nil { + proxyURL = account.Proxy.URL() + } + + upstreamStart := time.Now() + resp, err := s.httpUpstream.Do(upstreamReq, proxyURL, account.ID, account.Concurrency) + SetOpsLatencyMs(c, OpsUpstreamLatencyMsKey, time.Since(upstreamStart).Milliseconds()) + if err != nil { + return nil, s.handleOpenAIUpstreamTransportError(ctx, c, account, err, false) + } + defer func() { _ = resp.Body.Close() }() + + if resp.StatusCode >= 400 { + respBody := s.readUpstreamErrorBody(resp) + resp.Body = io.NopCloser(bytes.NewReader(respBody)) + s.updateGrokUsageSnapshot(ctx, account.ID, xai.ParseQuotaHeaders(resp.Header, resp.StatusCode)) + upstreamMsg := sanitizeUpstreamErrorMessage(extractUpstreamErrorMessage(respBody)) + if upstreamMsg == "" { + upstreamMsg = fmt.Sprintf("xAI upstream returned status %d", resp.StatusCode) + } + appendOpsUpstreamError(c, OpsUpstreamErrorEvent{ + Platform: account.Platform, + AccountID: account.ID, + AccountName: account.Name, + UpstreamStatusCode: resp.StatusCode, + UpstreamRequestID: firstNonEmpty(resp.Header.Get("x-request-id"), resp.Header.Get("xai-request-id")), + Kind: "failover", + Message: upstreamMsg, + }) + s.handleGrokAccountUpstreamError(ctx, account, resp.StatusCode, resp.Header, respBody) + if s.shouldFailoverUpstreamError(resp.StatusCode) { + return nil, &UpstreamFailoverError{ + StatusCode: resp.StatusCode, + ResponseBody: respBody, + RetryableOnSameAccount: account.IsPoolMode() && account.IsPoolModeRetryableStatus(resp.StatusCode), + } + } + return s.handleErrorResponse(ctx, resp, c, account, patchedBody, upstreamModel) + } + + s.updateGrokUsageSnapshot(ctx, account.ID, xai.ParseQuotaHeaders(resp.Header, resp.StatusCode)) + + var usage *OpenAIUsage + var firstTokenMs *int + responseID := "" + if reqStream { + streamResult, err := s.handleStreamingResponse(ctx, resp, c, account, startTime, originalModel, upstreamModel) + if err != nil { + return nil, err + } + usage = streamResult.usage + firstTokenMs = streamResult.firstTokenMs + responseID = strings.TrimSpace(streamResult.responseID) + } else { + nonStreamResult, err := s.handleNonStreamingResponse(ctx, resp, c, account, originalModel, upstreamModel) + if err != nil { + return nil, err + } + usage = nonStreamResult.usage + responseID = strings.TrimSpace(nonStreamResult.responseID) + } + + if usage == nil { + usage = &OpenAIUsage{} + } + return &OpenAIForwardResult{ + RequestID: firstNonEmpty(resp.Header.Get("x-request-id"), resp.Header.Get("xai-request-id")), + ResponseID: responseID, + Usage: *usage, + Model: originalModel, + UpstreamModel: upstreamModel, + ReasoningEffort: ptrStringOrNil(normalizeOpenAIReasoningEffort(gjson.GetBytes(patchedBody, "reasoning.effort").String())), + Stream: reqStream, + OpenAIWSMode: false, + ResponseHeaders: resp.Header.Clone(), + Duration: time.Since(startTime), + FirstTokenMs: firstTokenMs, + }, nil +} + +func patchGrokResponsesBody(body []byte, upstreamModel string) ([]byte, error) { + if !json.Valid(body) { + return nil, fmt.Errorf("invalid json request body") + } + out, err := sjson.SetBytes(body, "model", upstreamModel) + if err != nil { + return nil, err + } + for _, unsupportedField := range []string{"prompt_cache_retention", "safety_identifier"} { + if gjson.GetBytes(out, unsupportedField).Exists() { + out, err = sjson.DeleteBytes(out, unsupportedField) + if err != nil { + return nil, err + } + } + } + return out, nil +} + +func buildGrokResponsesRequest(ctx context.Context, c *gin.Context, account *Account, body []byte, token string) (*http.Request, error) { + targetURL, err := xai.BuildResponsesURL(account.GetGrokBaseURL()) + if err != nil { + return nil, err + } + req, err := http.NewRequestWithContext(ctx, http.MethodPost, targetURL, bytes.NewReader(body)) + if err != nil { + return nil, err + } + req.Header.Set("Authorization", "Bearer "+token) + req.Header.Set("Content-Type", "application/json") + req.Header.Set("Accept", "application/json, text/event-stream") + req.Header.Set("User-Agent", "sub2api-grok/1.0") + if c != nil { + if v := c.GetHeader("OpenAI-Beta"); strings.TrimSpace(v) != "" { + req.Header.Set("OpenAI-Beta", v) + } + } + return req, nil +} + +func (s *OpenAIGatewayService) updateGrokUsageSnapshot(ctx context.Context, accountID int64, snapshot *xai.QuotaSnapshot) { + if s == nil || s.accountRepo == nil || accountID <= 0 || snapshot == nil { + return + } + if s.codexSnapshotThrottle != nil && !s.codexSnapshotThrottle.Allow(accountID, time.Now()) { + return + } + _ = s.accountRepo.UpdateExtra(ctx, accountID, map[string]any{ + grokQuotaSnapshotExtraKey: snapshot, + }) +} + +func (s *OpenAIGatewayService) handleGrokAccountUpstreamError(ctx context.Context, account *Account, statusCode int, headers http.Header, responseBody []byte) { + if s == nil || account == nil { + return + } + switch statusCode { + case http.StatusUnauthorized: + s.tempUnscheduleGrok(ctx, account, 10*time.Minute, "grok oauth token unauthorized") + case http.StatusForbidden: + s.tempUnscheduleGrok(ctx, account, 30*time.Minute, "grok entitlement or subscription tier denied") + case http.StatusTooManyRequests: + cooldown := 2 * time.Minute + if snapshot := xai.ParseQuotaHeaders(headers, statusCode); snapshot != nil && snapshot.RetryAfterSeconds != nil && *snapshot.RetryAfterSeconds > 0 { + cooldown = time.Duration(*snapshot.RetryAfterSeconds) * time.Second + } + s.tempUnscheduleGrok(ctx, account, cooldown, "grok rate limited") + default: + if statusCode >= 500 { + s.tempUnscheduleGrok(ctx, account, 2*time.Minute, "grok upstream temporary error") + } + } + _ = responseBody +} + +func (s *OpenAIGatewayService) tempUnscheduleGrok(ctx context.Context, account *Account, cooldown time.Duration, reason string) { + if s == nil || account == nil { + return + } + until := time.Now().Add(cooldown) + if account.TempUnschedulableUntil != nil && account.TempUnschedulableUntil.After(until) { + until = *account.TempUnschedulableUntil + } + s.BlockAccountScheduling(account, until, reason) + if s.accountRepo != nil { + stateCtx, cancel := openAIAccountStateContext(ctx) + defer cancel() + _ = s.accountRepo.SetTempUnschedulable(stateCtx, account.ID, until, reason) + } +} + +func ptrStringOrNil(value string) *string { + if strings.TrimSpace(value) == "" { + return nil + } + return &value +} diff --git a/backend/internal/service/openai_gateway_grok_test.go b/backend/internal/service/openai_gateway_grok_test.go new file mode 100644 index 0000000000..15dd6a3acb --- /dev/null +++ b/backend/internal/service/openai_gateway_grok_test.go @@ -0,0 +1,347 @@ +//go:build unit + +package service + +import ( + "bytes" + "context" + "encoding/json" + "io" + "net/http" + "net/http/httptest" + "strings" + "testing" + "time" + + "github.com/Wei-Shaw/sub2api/internal/pkg/xai" + "github.com/gin-gonic/gin" + "github.com/stretchr/testify/require" + "github.com/tidwall/gjson" +) + +func TestPatchGrokResponsesBodySetsMappedModelAndDropsUnsupportedFields(t *testing.T) { + t.Parallel() + + body := []byte(`{ + "model": "grok", + "input": "hello", + "prompt_cache_retention": "24h", + "safety_identifier": "user-1", + "reasoning": {"effort": "high"} + }`) + + patched, err := patchGrokResponsesBody(body, "grok-4.3") + require.NoError(t, err) + require.True(t, json.Valid(patched)) + require.Equal(t, "grok-4.3", gjson.GetBytes(patched, "model").String()) + require.False(t, gjson.GetBytes(patched, "prompt_cache_retention").Exists()) + require.False(t, gjson.GetBytes(patched, "safety_identifier").Exists()) + require.Equal(t, "high", gjson.GetBytes(patched, "reasoning.effort").String()) +} + +func TestBuildGrokResponsesRequestUsesAccountBaseURLAndBearerToken(t *testing.T) { + t.Setenv(xai.EnvAllowUnsafeURLOverrides, "true") + + account := &Account{ + Platform: PlatformGrok, + Type: AccountTypeOAuth, + Credentials: map[string]any{ + "base_url": "https://xai.test/v1/", + }, + } + + req, err := buildGrokResponsesRequest(context.Background(), nil, account, []byte(`{"model":"grok-4.3"}`), "access-token") + require.NoError(t, err) + require.Equal(t, http.MethodPost, req.Method) + require.Equal(t, "https://xai.test/v1/responses", req.URL.String()) + require.Equal(t, "Bearer access-token", req.Header.Get("Authorization")) + require.Equal(t, "application/json", req.Header.Get("Content-Type")) + require.Contains(t, req.Header.Get("Accept"), "text/event-stream") + + data, err := io.ReadAll(req.Body) + require.NoError(t, err) + require.Equal(t, `{"model":"grok-4.3"}`, strings.TrimSpace(string(data))) +} + +func TestBuildGrokResponsesRequestRejectsUnsafeAccountBaseURL(t *testing.T) { + t.Parallel() + + account := &Account{ + Platform: PlatformGrok, + Type: AccountTypeOAuth, + Credentials: map[string]any{ + "base_url": "https://xai.test/v1", + }, + } + + _, err := buildGrokResponsesRequest(context.Background(), nil, account, []byte(`{"model":"grok-4.3"}`), "access-token") + require.Error(t, err) + require.Contains(t, err.Error(), "invalid base url") +} + +func TestForwardAsChatCompletionsForGrokUsesXAIChatCompletionsAndSnapshots(t *testing.T) { + gin.SetMode(gin.TestMode) + + recorder := httptest.NewRecorder() + c, _ := gin.CreateTestContext(recorder) + body := []byte(`{"model":"grok","messages":[{"role":"user","content":"hi"}],"stream":false}`) + c.Request = httptest.NewRequest(http.MethodPost, "/v1/chat/completions", bytes.NewReader(body)) + + account := &Account{ + ID: 51, + Name: "grok", + Platform: PlatformGrok, + Type: AccountTypeOAuth, + Concurrency: 1, + Credentials: map[string]any{ + "access_token": "access-token", + "expires_at": time.Now().Add(time.Hour).UTC().Format(time.RFC3339), + "base_url": xai.DefaultCLIBaseURL, + }, + } + repo := &grokQuotaAccountRepo{ + mockAccountRepoForPlatform: &mockAccountRepoForPlatform{ + accountsByID: map[int64]*Account{51: account}, + }, + } + upstream := &httpUpstreamRecorder{resp: &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{ + "Content-Type": []string{"application/json"}, + "Xai-Request-Id": []string{"xai-req"}, + "X-Ratelimit-Limit-Requests": []string{"10"}, + "X-Ratelimit-Remaining-Requests": []string{"9"}, + "X-Ratelimit-Limit-Tokens": []string{"1000"}, + "X-Ratelimit-Remaining-Tokens": []string{"990"}, + }, + Body: io.NopCloser(strings.NewReader(`{"id":"chatcmpl","object":"chat.completion","model":"grok-4.3","choices":[],"usage":{"prompt_tokens":1,"completion_tokens":2}}`)), + }} + svc := &OpenAIGatewayService{ + httpUpstream: upstream, + grokTokenProvider: NewGrokTokenProvider(repo, nil), + accountRepo: repo, + } + + result, err := svc.ForwardAsChatCompletions(context.Background(), c, account, body, "", "") + require.NoError(t, err) + require.Equal(t, xai.DefaultCLIBaseURL+"/chat/completions", upstream.lastReq.URL.String()) + require.Equal(t, "Bearer access-token", upstream.lastReq.Header.Get("Authorization")) + require.Equal(t, "grok-4.3", gjson.GetBytes(upstream.lastBody, "model").String()) + require.Equal(t, "grok", result.Model) + require.Equal(t, "grok-4.3", result.UpstreamModel) + require.Equal(t, 1, result.Usage.InputTokens) + require.Equal(t, 2, result.Usage.OutputTokens) + require.NotNil(t, repo.updates[51][grokQuotaSnapshotExtraKey]) + require.Equal(t, http.StatusOK, recorder.Code) +} + +func TestForwardGrokResponsesStreamingUsesXAIResponsesAndSnapshots(t *testing.T) { + gin.SetMode(gin.TestMode) + + recorder := httptest.NewRecorder() + c, _ := gin.CreateTestContext(recorder) + body := []byte(`{"model":"grok","input":"hi","stream":true}`) + c.Request = httptest.NewRequest(http.MethodPost, "/v1/responses", bytes.NewReader(body)) + c.Request.Header.Set("Content-Type", "application/json") + c.Request.Header.Set("OpenAI-Beta", "responses=experimental") + + account := &Account{ + ID: 52, + Name: "grok", + Platform: PlatformGrok, + Type: AccountTypeOAuth, + Concurrency: 1, + Credentials: map[string]any{ + "access_token": "access-token", + "expires_at": time.Now().Add(time.Hour).UTC().Format(time.RFC3339), + "base_url": xai.DefaultCLIBaseURL, + }, + } + repo := &grokQuotaAccountRepo{ + mockAccountRepoForPlatform: &mockAccountRepoForPlatform{ + accountsByID: map[int64]*Account{52: account}, + }, + } + upstreamBody := strings.Join([]string{ + `data: {"type":"response.output_text.delta","sequence_number":0,"delta":"ok"}`, + "", + `data: {"type":"response.completed","sequence_number":1,"response":{"id":"resp_grok","model":"grok-4.3","usage":{"input_tokens":5,"output_tokens":3,"input_tokens_details":{"cached_tokens":2}}}}`, + "", + }, "\n") + upstream := &httpUpstreamRecorder{resp: &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{ + "Content-Type": []string{"text/event-stream"}, + "Xai-Request-Id": []string{"xai-stream-req"}, + "X-Ratelimit-Limit-Requests": []string{"10"}, + "X-Ratelimit-Remaining-Requests": []string{"8"}, + "X-Ratelimit-Limit-Tokens": []string{"1000"}, + "X-Ratelimit-Remaining-Tokens": []string{"990"}, + }, + Body: io.NopCloser(strings.NewReader(upstreamBody)), + }} + svc := &OpenAIGatewayService{ + httpUpstream: upstream, + grokTokenProvider: NewGrokTokenProvider(repo, nil), + accountRepo: repo, + } + + result, err := svc.forwardGrokResponses(context.Background(), c, account, body, "grok", true, time.Now()) + require.NoError(t, err) + require.Equal(t, xai.DefaultCLIBaseURL+"/responses", upstream.lastReq.URL.String()) + require.Equal(t, "Bearer access-token", upstream.lastReq.Header.Get("Authorization")) + require.Equal(t, "responses=experimental", upstream.lastReq.Header.Get("OpenAI-Beta")) + require.Equal(t, "grok-4.3", gjson.GetBytes(upstream.lastBody, "model").String()) + require.True(t, gjson.GetBytes(upstream.lastBody, "stream").Bool()) + require.True(t, result.Stream) + require.Equal(t, "resp_grok", result.ResponseID) + require.Equal(t, "xai-stream-req", result.RequestID) + require.Equal(t, 5, result.Usage.InputTokens) + require.Equal(t, 3, result.Usage.OutputTokens) + require.Equal(t, 2, result.Usage.CacheReadInputTokens) + require.Contains(t, recorder.Header().Get("Content-Type"), "text/event-stream") + require.Contains(t, recorder.Body.String(), "response.output_text.delta") + require.NotNil(t, repo.updates[52][grokQuotaSnapshotExtraKey]) +} + +func TestForwardAsChatCompletionsForGrokStreamingUsesRawXAIChatCompletions(t *testing.T) { + gin.SetMode(gin.TestMode) + + recorder := httptest.NewRecorder() + c, _ := gin.CreateTestContext(recorder) + body := []byte(`{"model":"grok","messages":[{"role":"user","content":"hi"}],"stream":true}`) + c.Request = httptest.NewRequest(http.MethodPost, "/v1/chat/completions", bytes.NewReader(body)) + c.Request.Header.Set("Content-Type", "application/json") + + account := &Account{ + ID: 53, + Name: "grok", + Platform: PlatformGrok, + Type: AccountTypeOAuth, + Concurrency: 1, + Credentials: map[string]any{ + "access_token": "access-token", + "expires_at": time.Now().Add(time.Hour).UTC().Format(time.RFC3339), + "base_url": xai.DefaultCLIBaseURL, + }, + } + repo := &grokQuotaAccountRepo{ + mockAccountRepoForPlatform: &mockAccountRepoForPlatform{ + accountsByID: map[int64]*Account{53: account}, + }, + } + upstreamBody := strings.Join([]string{ + `data: {"id":"chatcmpl_grok","object":"chat.completion.chunk","model":"grok-4.3","choices":[{"index":0,"delta":{"content":"ok"}}]}`, + "", + `data: {"id":"chatcmpl_grok","object":"chat.completion.chunk","model":"grok-4.3","choices":[],"usage":{"prompt_tokens":6,"completion_tokens":4,"total_tokens":10,"prompt_tokens_details":{"cached_tokens":1}}}`, + "", + "data: [DONE]", + "", + }, "\n") + upstream := &httpUpstreamRecorder{resp: &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{ + "Content-Type": []string{"text/event-stream"}, + "X-Request-Id": []string{"chat-stream-req"}, + "X-Ratelimit-Limit-Requests": []string{"10"}, + "X-Ratelimit-Remaining-Requests": []string{"7"}, + }, + Body: io.NopCloser(strings.NewReader(upstreamBody)), + }} + svc := &OpenAIGatewayService{ + cfg: rawChatCompletionsTestConfig(), + httpUpstream: upstream, + grokTokenProvider: NewGrokTokenProvider(repo, nil), + accountRepo: repo, + } + + result, err := svc.ForwardAsChatCompletions(context.Background(), c, account, body, "", "") + require.NoError(t, err) + require.Equal(t, xai.DefaultCLIBaseURL+"/chat/completions", upstream.lastReq.URL.String()) + require.Equal(t, "Bearer access-token", upstream.lastReq.Header.Get("Authorization")) + require.Equal(t, "text/event-stream", upstream.lastReq.Header.Get("Accept")) + require.Equal(t, "sub2api-grok/1.0", upstream.lastReq.Header.Get("User-Agent")) + require.Equal(t, "grok-4.3", gjson.GetBytes(upstream.lastBody, "model").String()) + require.True(t, gjson.GetBytes(upstream.lastBody, "stream_options.include_usage").Bool()) + require.True(t, result.Stream) + require.Equal(t, 6, result.Usage.InputTokens) + require.Equal(t, 4, result.Usage.OutputTokens) + require.Equal(t, 1, result.Usage.CacheReadInputTokens) + require.Contains(t, recorder.Body.String(), "data: [DONE]") + require.NotNil(t, repo.updates[53][grokQuotaSnapshotExtraKey]) +} + +func TestHandleGrokAccountUpstreamErrorTempUnschedulesReadinessStates(t *testing.T) { + tests := []struct { + name string + status int + headers http.Header + wantReason string + wantMinCooldown time.Duration + wantMaxCooldown time.Duration + }{ + { + name: "unauthorized reauth", + status: http.StatusUnauthorized, + wantReason: "grok oauth token unauthorized", + wantMinCooldown: 10*time.Minute - time.Second, + wantMaxCooldown: 10*time.Minute + time.Second, + }, + { + name: "forbidden entitlement", + status: http.StatusForbidden, + wantReason: "grok entitlement or subscription tier denied", + wantMinCooldown: 30*time.Minute - time.Second, + wantMaxCooldown: 30*time.Minute + time.Second, + }, + { + name: "rate limited retry after", + status: http.StatusTooManyRequests, + headers: http.Header{"Retry-After": []string{"45"}}, + wantReason: "grok rate limited", + wantMinCooldown: 44 * time.Second, + wantMaxCooldown: 46 * time.Second, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + account := &Account{ID: 61, Platform: PlatformGrok, Type: AccountTypeOAuth} + repo := &grokQuotaAccountRepo{} + svc := &OpenAIGatewayService{accountRepo: repo} + before := time.Now() + + svc.handleGrokAccountUpstreamError(context.Background(), account, tt.status, tt.headers, nil) + + require.True(t, svc.isOpenAIAccountRuntimeBlocked(account)) + require.Equal(t, 1, repo.tempUnschedCalls) + require.Equal(t, account.ID, repo.lastTempUnschedID) + require.Equal(t, tt.wantReason, repo.lastTempUnschedReason) + require.True(t, repo.lastTempUnschedUntil.After(before.Add(tt.wantMinCooldown))) + require.True(t, repo.lastTempUnschedUntil.Before(before.Add(tt.wantMaxCooldown))) + }) + } +} + +func TestHandleGrokAccountUpstreamErrorDoesNotShortenExistingPause(t *testing.T) { + existingUntil := time.Now().Add(15 * time.Minute) + account := &Account{ + ID: 62, + Platform: PlatformGrok, + Type: AccountTypeOAuth, + TempUnschedulableUntil: &existingUntil, + TempUnschedulableReason: "existing pause", + } + repo := &grokQuotaAccountRepo{} + svc := &OpenAIGatewayService{accountRepo: repo} + + svc.handleGrokAccountUpstreamError(context.Background(), account, http.StatusTooManyRequests, http.Header{"Retry-After": []string{"45"}}, nil) + + require.Equal(t, 1, repo.tempUnschedCalls) + require.WithinDuration(t, existingUntil, repo.lastTempUnschedUntil, time.Second) + value, ok := svc.openaiAccountRuntimeBlockUntil.Load(account.ID) + require.True(t, ok) + runtimeUntil, ok := value.(time.Time) + require.True(t, ok) + require.WithinDuration(t, existingUntil, runtimeUntil, time.Second) +} diff --git a/backend/internal/service/openai_gateway_record_usage_test.go b/backend/internal/service/openai_gateway_record_usage_test.go index 1e4e58dc42..83ca42b733 100644 --- a/backend/internal/service/openai_gateway_record_usage_test.go +++ b/backend/internal/service/openai_gateway_record_usage_test.go @@ -224,6 +224,7 @@ func newOpenAIRecordUsageServiceForTest(usageRepo UsageLogRepository, userRepo U nil, nil, nil, + nil, nil, // userPlatformQuotaRepo ) svc.userGroupRateResolver = newUserGroupRateResolver( diff --git a/backend/internal/service/openai_gateway_service.go b/backend/internal/service/openai_gateway_service.go index 3fdad536d0..2758a3b0a1 100644 --- a/backend/internal/service/openai_gateway_service.go +++ b/backend/internal/service/openai_gateway_service.go @@ -26,6 +26,7 @@ import ( "github.com/Wei-Shaw/sub2api/internal/pkg/logger" "github.com/Wei-Shaw/sub2api/internal/pkg/openai" "github.com/Wei-Shaw/sub2api/internal/pkg/openai_compat" + "github.com/Wei-Shaw/sub2api/internal/pkg/xai" "github.com/Wei-Shaw/sub2api/internal/util/responseheaders" "github.com/Wei-Shaw/sub2api/internal/util/urlvalidator" "github.com/cespare/xxhash/v2" @@ -349,6 +350,7 @@ type OpenAIGatewayService struct { httpUpstream HTTPUpstream deferredService *DeferredService openAITokenProvider *OpenAITokenProvider + grokTokenProvider *GrokTokenProvider toolCorrector *CodexToolCorrector openaiWSResolver OpenAIWSProtocolResolver resolver *ModelPricingResolver @@ -396,6 +398,7 @@ func NewOpenAIGatewayService( httpUpstream HTTPUpstream, deferredService *DeferredService, openAITokenProvider *OpenAITokenProvider, + grokTokenProvider *GrokTokenProvider, resolver *ModelPricingResolver, channelService *ChannelService, balanceNotifyService *BalanceNotifyService, @@ -426,6 +429,7 @@ func NewOpenAIGatewayService( httpUpstream: httpUpstream, deferredService: deferredService, openAITokenProvider: openAITokenProvider, + grokTokenProvider: grokTokenProvider, toolCorrector: NewCodexToolCorrector(), openaiWSResolver: NewOpenAIWSProtocolResolver(cfg), resolver: resolver, @@ -1317,11 +1321,18 @@ func (s *OpenAIGatewayService) SelectAccountForModel(ctx context.Context, groupI // SelectAccountForModelWithExclusions selects an account supporting the requested model while excluding specified accounts. // SelectAccountForModelWithExclusions 选择支持指定模型的账号,同时排除指定的账号。 func (s *OpenAIGatewayService) SelectAccountForModelWithExclusions(ctx context.Context, groupID *int64, sessionHash string, requestedModel string, excludedIDs map[int64]struct{}) (*Account, error) { - return s.selectAccountForModelWithExclusions(s.withOpenAIQuotaAutoPauseContext(ctx), groupID, sessionHash, requestedModel, excludedIDs, false, 0, "") + return s.selectAccountForModelWithExclusions(s.withOpenAIQuotaAutoPauseContext(ctx), groupID, PlatformOpenAI, sessionHash, requestedModel, excludedIDs, false, 0, "") } // noAvailableOpenAISelectionError builds the standard "no account available" error // while preserving the compact-specific error when applicable. +func normalizeOpenAICompatiblePlatform(platform string) string { + if platform == PlatformGrok { + return PlatformGrok + } + return PlatformOpenAI +} + func noAvailableOpenAISelectionError(requestedModel string, compactBlocked bool) error { if compactBlocked { return ErrNoAvailableCompactAccounts @@ -1348,22 +1359,34 @@ func openAICompactSupportTier(account *Account) int { return 0 } -// isOpenAIAccountEligibleForRequest centralises the schedulable / OpenAI / model / -// compact-support checks used during account selection. -func isOpenAIAccountEligibleForRequest(ctx context.Context, account *Account, requestedModel string, requireCompact bool, requiredCapability OpenAIEndpointCapability) bool { - if account == nil || !account.IsOpenAI() || !account.IsSchedulableForModelWithContext(ctx, requestedModel) { +func isOpenAICompatibleAccountEligibleForRequest(ctx context.Context, account *Account, platform string, requestedModel string, requireCompact bool, requiredCapability OpenAIEndpointCapability) bool { + platform = normalizeOpenAICompatiblePlatform(platform) + if account == nil || account.Platform != platform || !account.IsOpenAICompatible() || !account.IsSchedulableForModelWithContext(ctx, requestedModel) { return false } - if paused, reason := shouldAutoPauseOpenAIAccountByQuota(ctx, account); paused { - // Debug level: this fires per-candidate on the scheduling hot path, so Info - // would amplify into log spam once several accounts cross the threshold. - slog.Debug("account_auto_paused_by_quota", - "account_id", account.ID, - "window", reason.window, - "threshold", reason.threshold, - "utilization", reason.utilization, - ) - return false + if account.IsOpenAI() { + if paused, reason := shouldAutoPauseOpenAIAccountByQuota(ctx, account); paused { + // Debug level: this fires per-candidate on the scheduling hot path, so Info + // would amplify into log spam once several accounts cross the threshold. + slog.Debug("account_auto_paused_by_quota", + "account_id", account.ID, + "window", reason.window, + "threshold", reason.threshold, + "utilization", reason.utilization, + ) + return false + } + } + if account.IsGrok() { + if paused, reason := shouldAutoPauseGrokAccountByQuota(account); paused { + slog.Debug("grok_account_auto_paused_by_quota", + "account_id", account.ID, + "window", reason.window, + "threshold", reason.threshold, + "utilization", reason.utilization, + ) + return false + } } if requestedModel != "" && !account.IsModelSupported(requestedModel) { return false @@ -1371,7 +1394,7 @@ func isOpenAIAccountEligibleForRequest(ctx context.Context, account *Account, re if !account.SupportsOpenAIEndpointCapability(requiredCapability) { return false } - if requireCompact && openAICompactSupportTier(account) == 0 { + if requireCompact && (!account.IsOpenAI() || openAICompactSupportTier(account) == 0) { return false } return true @@ -1383,6 +1406,70 @@ type openAIQuotaAutoPauseDecision struct { utilization float64 } +func shouldAutoPauseGrokAccountByQuota(account *Account) (bool, openAIQuotaAutoPauseDecision) { + if account == nil || !account.IsGrok() || account.Type != AccountTypeOAuth { + return false, openAIQuotaAutoPauseDecision{} + } + snapshot, err := grokQuotaSnapshotFromExtra(account.Extra) + if err != nil || snapshot == nil { + return false, openAIQuotaAutoPauseDecision{} + } + now := time.Now() + if grokQuotaSnapshotStaleForPause(snapshot, now) { + return false, openAIQuotaAutoPauseDecision{} + } + if grokQuotaRetryAfterActive(snapshot, now) { + return true, openAIQuotaAutoPauseDecision{window: "retry_after", threshold: 1, utilization: 1} + } + if paused, decision := shouldAutoPauseGrokQuotaWindow("requests", snapshot.Requests, now); paused { + return true, decision + } + if paused, decision := shouldAutoPauseGrokQuotaWindow("tokens", snapshot.Tokens, now); paused { + return true, decision + } + return false, openAIQuotaAutoPauseDecision{} +} + +func grokQuotaRetryAfterActive(snapshot *xai.QuotaSnapshot, now time.Time) bool { + if snapshot == nil || snapshot.RetryAfterSeconds == nil || *snapshot.RetryAfterSeconds <= 0 { + return false + } + if strings.TrimSpace(snapshot.UpdatedAt) == "" { + return true + } + updatedAt, err := parseTime(snapshot.UpdatedAt) + if err != nil { + return true + } + retryAfterUntil := updatedAt.Add(time.Duration(*snapshot.RetryAfterSeconds) * time.Second) + return now.Before(retryAfterUntil) +} + +func shouldAutoPauseGrokQuotaWindow(name string, window *xai.QuotaWindow, now time.Time) (bool, openAIQuotaAutoPauseDecision) { + if window == nil || window.Limit == nil || window.Remaining == nil || *window.Limit <= 0 { + return false, openAIQuotaAutoPauseDecision{} + } + if window.ResetUnix != nil && *window.ResetUnix > 0 && !now.Before(time.Unix(*window.ResetUnix, 0)) { + return false, openAIQuotaAutoPauseDecision{} + } + utilization := float64(*window.Limit-*window.Remaining) / float64(*window.Limit) + if *window.Remaining <= 0 || utilization >= 1 { + return true, openAIQuotaAutoPauseDecision{window: name, threshold: 1, utilization: utilization} + } + return false, openAIQuotaAutoPauseDecision{} +} + +func grokQuotaSnapshotStaleForPause(snapshot *xai.QuotaSnapshot, now time.Time) bool { + if snapshot == nil || strings.TrimSpace(snapshot.UpdatedAt) == "" { + return false + } + updatedAt, err := parseTime(snapshot.UpdatedAt) + if err != nil { + return false + } + return now.Sub(updatedAt) >= openAICodexAutoPauseStaleAfter +} + func shouldAutoPauseOpenAIAccountByQuota(ctx context.Context, account *Account) (bool, openAIQuotaAutoPauseDecision) { if account == nil || !account.IsOpenAI() { return false, openAIQuotaAutoPauseDecision{} @@ -1640,7 +1727,8 @@ func resolveOpenAIAccountUpstreamModelForRequest(account *Account, requestedMode return upstreamModel } -func (s *OpenAIGatewayService) selectAccountForModelWithExclusions(ctx context.Context, groupID *int64, sessionHash string, requestedModel string, excludedIDs map[int64]struct{}, requireCompact bool, stickyAccountID int64, requiredCapability OpenAIEndpointCapability) (*Account, error) { +func (s *OpenAIGatewayService) selectAccountForModelWithExclusions(ctx context.Context, groupID *int64, platform string, sessionHash string, requestedModel string, excludedIDs map[int64]struct{}, requireCompact bool, stickyAccountID int64, requiredCapability OpenAIEndpointCapability) (*Account, error) { + platform = normalizeOpenAICompatiblePlatform(platform) if s.checkChannelPricingRestriction(ctx, groupID, requestedModel) { slog.Warn("channel pricing restriction blocked request", "group_id", derefGroupID(groupID), @@ -1650,20 +1738,20 @@ func (s *OpenAIGatewayService) selectAccountForModelWithExclusions(ctx context.C // 1. 尝试粘性会话命中 // Try sticky session hit - if account := s.tryStickySessionHit(ctx, groupID, sessionHash, requestedModel, excludedIDs, requireCompact, stickyAccountID, requiredCapability); account != nil { + if account := s.tryStickySessionHit(ctx, groupID, platform, sessionHash, requestedModel, excludedIDs, requireCompact, stickyAccountID, requiredCapability); account != nil { return account, nil } // 2. 获取可调度的 OpenAI 账号 // Get schedulable OpenAI accounts - accounts, err := s.listSchedulableAccounts(ctx, groupID) + accounts, err := s.listSchedulableAccounts(ctx, groupID, platform) if err != nil { return nil, fmt.Errorf("query accounts failed: %w", err) } // 3. 按优先级 + LRU 选择最佳账号 // Select by priority + LRU - selected, compactBlocked := s.selectBestAccount(ctx, groupID, accounts, requestedModel, excludedIDs, requireCompact, requiredCapability) + selected, compactBlocked := s.selectBestAccount(ctx, groupID, platform, accounts, requestedModel, excludedIDs, requireCompact, requiredCapability) if selected == nil { return nil, noAvailableOpenAISelectionError(requestedModel, compactBlocked) @@ -1688,10 +1776,11 @@ func (s *OpenAIGatewayService) selectAccountForModelWithExclusions(ctx context.C // // tryStickySessionHit attempts to get account from sticky session. // Returns account if hit and usable; clears session and returns nil if account is unavailable. -func (s *OpenAIGatewayService) tryStickySessionHit(ctx context.Context, groupID *int64, sessionHash, requestedModel string, excludedIDs map[int64]struct{}, requireCompact bool, stickyAccountID int64, requiredCapability OpenAIEndpointCapability) *Account { +func (s *OpenAIGatewayService) tryStickySessionHit(ctx context.Context, groupID *int64, platform string, sessionHash, requestedModel string, excludedIDs map[int64]struct{}, requireCompact bool, stickyAccountID int64, requiredCapability OpenAIEndpointCapability) *Account { if sessionHash == "" { return nil } + platform = normalizeOpenAICompatiblePlatform(platform) accountID := stickyAccountID if accountID <= 0 { @@ -1720,14 +1809,14 @@ func (s *OpenAIGatewayService) tryStickySessionHit(ctx context.Context, groupID // 验证账号是否可用于当前请求 // Verify account is usable for current request - if !isOpenAIAccountEligibleForRequest(ctx, account, requestedModel, false, requiredCapability) { + if !isOpenAICompatibleAccountEligibleForRequest(ctx, account, platform, requestedModel, false, requiredCapability) { return nil } if s.isOpenAIAccountRuntimeBlocked(account) { _ = s.deleteStickySessionAccountID(ctx, groupID, sessionHash) return nil } - account = s.recheckSelectedOpenAIAccountFromDB(ctx, account, requestedModel, requireCompact, requiredCapability) + account = s.recheckSelectedOpenAIAccountFromDB(ctx, account, platform, requestedModel, requireCompact, requiredCapability) if account == nil || !openAIStickyAccountMatchesGroup(account, groupID) { _ = s.deleteStickySessionAccountID(ctx, groupID, sessionHash) return nil @@ -1751,7 +1840,8 @@ func (s *OpenAIGatewayService) tryStickySessionHit(ctx context.Context, groupID // Returns nil if no available account. The second return reports whether at // least one candidate was filtered out solely because it lacks compact support // (only meaningful when requireCompact=true). -func (s *OpenAIGatewayService) selectBestAccount(ctx context.Context, groupID *int64, accounts []Account, requestedModel string, excludedIDs map[int64]struct{}, requireCompact bool, requiredCapability OpenAIEndpointCapability) (*Account, bool) { +func (s *OpenAIGatewayService) selectBestAccount(ctx context.Context, groupID *int64, platform string, accounts []Account, requestedModel string, excludedIDs map[int64]struct{}, requireCompact bool, requiredCapability OpenAIEndpointCapability) (*Account, bool) { + platform = normalizeOpenAICompatiblePlatform(platform) var selected *Account selectedCompactTier := -1 compactBlocked := false @@ -1766,11 +1856,11 @@ func (s *OpenAIGatewayService) selectBestAccount(ctx context.Context, groupID *i continue } - fresh := s.resolveFreshSchedulableOpenAIAccount(ctx, acc, requestedModel, false, requiredCapability) + fresh := s.resolveFreshSchedulableOpenAIAccount(ctx, acc, platform, requestedModel, false, requiredCapability) if fresh == nil { continue } - fresh = s.recheckSelectedOpenAIAccountFromDB(ctx, fresh, requestedModel, false, requiredCapability) + fresh = s.recheckSelectedOpenAIAccountFromDB(ctx, fresh, platform, requestedModel, false, requiredCapability) if fresh == nil { continue } @@ -1847,10 +1937,11 @@ func (s *OpenAIGatewayService) isBetterAccount(candidate, current *Account) bool // SelectAccountWithLoadAwareness selects an account with load-awareness and wait plan. func (s *OpenAIGatewayService) SelectAccountWithLoadAwareness(ctx context.Context, groupID *int64, sessionHash string, requestedModel string, excludedIDs map[int64]struct{}) (*AccountSelectionResult, error) { - return s.selectAccountWithLoadAwareness(s.withOpenAIQuotaAutoPauseContext(ctx), groupID, sessionHash, requestedModel, excludedIDs, false, "") + return s.selectAccountWithLoadAwareness(s.withOpenAIQuotaAutoPauseContext(ctx), groupID, PlatformOpenAI, sessionHash, requestedModel, excludedIDs, false, "") } -func (s *OpenAIGatewayService) selectAccountWithLoadAwareness(ctx context.Context, groupID *int64, sessionHash string, requestedModel string, excludedIDs map[int64]struct{}, requireCompact bool, requiredCapability OpenAIEndpointCapability) (*AccountSelectionResult, error) { +func (s *OpenAIGatewayService) selectAccountWithLoadAwareness(ctx context.Context, groupID *int64, platform string, sessionHash string, requestedModel string, excludedIDs map[int64]struct{}, requireCompact bool, requiredCapability OpenAIEndpointCapability) (*AccountSelectionResult, error) { + platform = normalizeOpenAICompatiblePlatform(platform) if s.checkChannelPricingRestriction(ctx, groupID, requestedModel) { slog.Warn("channel pricing restriction blocked request", "group_id", derefGroupID(groupID), @@ -1867,7 +1958,7 @@ func (s *OpenAIGatewayService) selectAccountWithLoadAwareness(ctx context.Contex } } if s.concurrencyService == nil || !cfg.LoadBatchEnabled { - account, err := s.selectAccountForModelWithExclusions(ctx, groupID, sessionHash, requestedModel, excludedIDs, requireCompact, stickyAccountID, requiredCapability) + account, err := s.selectAccountForModelWithExclusions(ctx, groupID, platform, sessionHash, requestedModel, excludedIDs, requireCompact, stickyAccountID, requiredCapability) if err != nil { return nil, err } @@ -1894,7 +1985,7 @@ func (s *OpenAIGatewayService) selectAccountWithLoadAwareness(ctx context.Contex }) } - accounts, err := s.listSchedulableAccounts(ctx, groupID) + accounts, err := s.listSchedulableAccounts(ctx, groupID, platform) if err != nil { return nil, err } @@ -1920,8 +2011,8 @@ func (s *OpenAIGatewayService) selectAccountWithLoadAwareness(ctx context.Contex if clearSticky { _ = s.deleteStickySessionAccountID(ctx, groupID, sessionHash) } - if !clearSticky && isOpenAIAccountEligibleForRequest(ctx, account, requestedModel, false, requiredCapability) { - account = s.recheckSelectedOpenAIAccountFromDB(ctx, account, requestedModel, requireCompact, requiredCapability) + if !clearSticky && isOpenAICompatibleAccountEligibleForRequest(ctx, account, platform, requestedModel, false, requiredCapability) { + account = s.recheckSelectedOpenAIAccountFromDB(ctx, account, platform, requestedModel, requireCompact, requiredCapability) if account == nil { _ = s.deleteStickySessionAccountID(ctx, groupID, sessionHash) } else if !openAIStickyAccountMatchesGroup(account, groupID) { @@ -1967,7 +2058,7 @@ func (s *OpenAIGatewayService) selectAccountWithLoadAwareness(ctx context.Contex // Scheduler snapshots can be temporarily stale (bucket rebuild is throttled); // re-check schedulability here so recently rate-limited/overloaded accounts // are not selected again before the bucket is rebuilt. - if !isOpenAIAccountEligibleForRequest(ctx, acc, requestedModel, false, requiredCapability) { + if !isOpenAICompatibleAccountEligibleForRequest(ctx, acc, platform, requestedModel, false, requiredCapability) { continue } if s.isOpenAIAccountRuntimeBlocked(acc) { @@ -2052,11 +2143,11 @@ func (s *OpenAIGatewayService) selectAccountWithLoadAwareness(ctx context.Contex } for _, item := range selectionOrder { - fresh := s.resolveFreshSchedulableOpenAIAccount(ctx, item.account, requestedModel, false, requiredCapability) + fresh := s.resolveFreshSchedulableOpenAIAccount(ctx, item.account, platform, requestedModel, false, requiredCapability) if fresh == nil { continue } - fresh = s.recheckSelectedOpenAIAccountFromDB(ctx, fresh, requestedModel, requireCompact, requiredCapability) + fresh = s.recheckSelectedOpenAIAccountFromDB(ctx, fresh, platform, requestedModel, requireCompact, requiredCapability) if fresh == nil { continue } @@ -2086,11 +2177,11 @@ func (s *OpenAIGatewayService) selectAccountWithLoadAwareness(ctx context.Contex ordered = prioritizeOpenAICompactAccounts(ordered) } for _, acc := range ordered { - fresh := s.resolveFreshSchedulableOpenAIAccount(ctx, acc, requestedModel, false, requiredCapability) + fresh := s.resolveFreshSchedulableOpenAIAccount(ctx, acc, platform, requestedModel, false, requiredCapability) if fresh == nil { continue } - fresh = s.recheckSelectedOpenAIAccountFromDB(ctx, fresh, requestedModel, requireCompact, requiredCapability) + fresh = s.recheckSelectedOpenAIAccountFromDB(ctx, fresh, platform, requestedModel, requireCompact, requiredCapability) if fresh == nil { continue } @@ -2131,11 +2222,11 @@ func (s *OpenAIGatewayService) selectAccountWithLoadAwareness(ctx context.Contex candidates = prioritizeOpenAICompactAccounts(candidates) } for _, acc := range candidates { - fresh := s.resolveFreshSchedulableOpenAIAccount(ctx, acc, requestedModel, false, requiredCapability) + fresh := s.resolveFreshSchedulableOpenAIAccount(ctx, acc, platform, requestedModel, false, requiredCapability) if fresh == nil { continue } - fresh = s.recheckSelectedOpenAIAccountFromDB(ctx, fresh, requestedModel, requireCompact, requiredCapability) + fresh = s.recheckSelectedOpenAIAccountFromDB(ctx, fresh, platform, requestedModel, requireCompact, requiredCapability) if fresh == nil { continue } @@ -2156,19 +2247,20 @@ func (s *OpenAIGatewayService) selectAccountWithLoadAwareness(ctx context.Contex return nil, ErrNoAvailableAccounts } -func (s *OpenAIGatewayService) listSchedulableAccounts(ctx context.Context, groupID *int64) ([]Account, error) { +func (s *OpenAIGatewayService) listSchedulableAccounts(ctx context.Context, groupID *int64, platform string) ([]Account, error) { + platform = normalizeOpenAICompatiblePlatform(platform) if s.schedulerSnapshot != nil { - accounts, _, err := s.schedulerSnapshot.ListSchedulableAccounts(ctx, groupID, PlatformOpenAI, false) + accounts, _, err := s.schedulerSnapshot.ListSchedulableAccounts(ctx, groupID, platform, false) return accounts, err } var accounts []Account var err error if s.cfg != nil && s.cfg.RunMode == config.RunModeSimple { - accounts, err = s.accountRepo.ListSchedulableByPlatform(ctx, PlatformOpenAI) + accounts, err = s.accountRepo.ListSchedulableByPlatform(ctx, platform) } else if groupID != nil { - accounts, err = s.accountRepo.ListSchedulableByGroupIDAndPlatform(ctx, *groupID, PlatformOpenAI) + accounts, err = s.accountRepo.ListSchedulableByGroupIDAndPlatform(ctx, *groupID, platform) } else { - accounts, err = s.accountRepo.ListSchedulableUngroupedByPlatform(ctx, PlatformOpenAI) + accounts, err = s.accountRepo.ListSchedulableUngroupedByPlatform(ctx, platform) } if err != nil { return nil, fmt.Errorf("query accounts failed: %w", err) @@ -2183,10 +2275,11 @@ func (s *OpenAIGatewayService) tryAcquireAccountSlot(ctx context.Context, accoun return s.concurrencyService.AcquireAccountSlot(ctx, accountID, maxConcurrency) } -func (s *OpenAIGatewayService) resolveFreshSchedulableOpenAIAccount(ctx context.Context, account *Account, requestedModel string, requireCompact bool, requiredCapability OpenAIEndpointCapability) *Account { +func (s *OpenAIGatewayService) resolveFreshSchedulableOpenAIAccount(ctx context.Context, account *Account, platform string, requestedModel string, requireCompact bool, requiredCapability OpenAIEndpointCapability) *Account { if account == nil { return nil } + platform = normalizeOpenAICompatiblePlatform(platform) fresh := account if s.schedulerSnapshot != nil { @@ -2197,7 +2290,7 @@ func (s *OpenAIGatewayService) resolveFreshSchedulableOpenAIAccount(ctx context. fresh = current } - if !isOpenAIAccountEligibleForRequest(ctx, fresh, requestedModel, requireCompact, requiredCapability) { + if !isOpenAICompatibleAccountEligibleForRequest(ctx, fresh, platform, requestedModel, requireCompact, requiredCapability) { return nil } if s.isOpenAIAccountRuntimeBlocked(fresh) { @@ -2206,12 +2299,13 @@ func (s *OpenAIGatewayService) resolveFreshSchedulableOpenAIAccount(ctx context. return fresh } -func (s *OpenAIGatewayService) recheckSelectedOpenAIAccountFromDB(ctx context.Context, account *Account, requestedModel string, requireCompact bool, requiredCapability OpenAIEndpointCapability) *Account { +func (s *OpenAIGatewayService) recheckSelectedOpenAIAccountFromDB(ctx context.Context, account *Account, platform string, requestedModel string, requireCompact bool, requiredCapability OpenAIEndpointCapability) *Account { if account == nil { return nil } + platform = normalizeOpenAICompatiblePlatform(platform) if s.schedulerSnapshot == nil || s.accountRepo == nil { - if !isOpenAIAccountEligibleForRequest(ctx, account, requestedModel, requireCompact, requiredCapability) { + if !isOpenAICompatibleAccountEligibleForRequest(ctx, account, platform, requestedModel, requireCompact, requiredCapability) { return nil } return account @@ -2221,7 +2315,7 @@ func (s *OpenAIGatewayService) recheckSelectedOpenAIAccountFromDB(ctx context.Co if err != nil || latest == nil { return nil } - if !isOpenAIAccountEligibleForRequest(ctx, latest, requestedModel, requireCompact, requiredCapability) { + if !isOpenAICompatibleAccountEligibleForRequest(ctx, latest, platform, requestedModel, requireCompact, requiredCapability) { return nil } if s.isOpenAIAccountRuntimeBlocked(latest) { @@ -2299,6 +2393,20 @@ func (s *OpenAIGatewayService) schedulingConfig() config.GatewaySchedulingConfig func (s *OpenAIGatewayService) GetAccessToken(ctx context.Context, account *Account) (string, string, error) { switch account.Type { case AccountTypeOAuth: + if account.Platform == PlatformGrok { + if s.grokTokenProvider != nil { + accessToken, err := s.grokTokenProvider.GetAccessToken(ctx, account) + if err != nil { + return "", "", err + } + return accessToken, "oauth", nil + } + accessToken := account.GetGrokAccessToken() + if accessToken == "" { + return "", "", errors.New("access_token not found in credentials") + } + return accessToken, "oauth", nil + } // 使用 TokenProvider 获取缓存的 token if s.openAITokenProvider != nil { accessToken, err := s.openAITokenProvider.GetAccessToken(ctx, account) @@ -2314,6 +2422,13 @@ func (s *OpenAIGatewayService) GetAccessToken(ctx context.Context, account *Acco } return accessToken, "oauth", nil case AccountTypeAPIKey: + if account.Platform == PlatformGrok { + apiKey := strings.TrimSpace(account.GetCredential("api_key")) + if apiKey == "" { + return "", "", errors.New("api_key not found in credentials") + } + return apiKey, "apikey", nil + } apiKey := account.GetOpenAIApiKey() if apiKey == "" { return "", "", errors.New("api_key not found in credentials") @@ -2405,6 +2520,11 @@ func (s *OpenAIGatewayService) Forward(ctx context.Context, c *gin.Context, acco reqModel, reqStream, promptCacheKey := requestView.Model, requestView.Stream, requestView.PromptCacheKey originalModel := reqModel + if account.Platform == PlatformGrok { + _ = promptCacheKey + return s.forwardGrokResponses(ctx, c, account, body, originalModel, reqStream, startTime) + } + if account.Type == AccountTypeAPIKey && !openai_compat.ShouldUseResponsesAPI(account.Extra) { return s.forwardResponsesViaRawChatCompletions(ctx, c, account, body) } diff --git a/backend/internal/service/openai_images_test.go b/backend/internal/service/openai_images_test.go index 74846e0a23..9897bffed0 100644 --- a/backend/internal/service/openai_images_test.go +++ b/backend/internal/service/openai_images_test.go @@ -442,6 +442,16 @@ func TestAccountSupportsOpenAIImageCapability_OAuthSupportsNative(t *testing.T) require.True(t, account.SupportsOpenAIImageCapability(OpenAIImagesCapabilityNative)) } +func TestAccountSupportsOpenAIImageCapability_EmptyRequirementDoesNotRejectGrok(t *testing.T) { + account := &Account{ + Platform: PlatformGrok, + Type: AccountTypeOAuth, + } + + require.True(t, account.SupportsOpenAIImageCapability("")) + require.False(t, account.SupportsOpenAIImageCapability(OpenAIImagesCapabilityBasic)) +} + func TestAccountSupportsOpenAIEndpointCapability(t *testing.T) { t.Run("OpenAI APIKey 默认兼容 chat 和 embeddings", func(t *testing.T) { account := &Account{ 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/openai_ws_protocol_forward_test.go b/backend/internal/service/openai_ws_protocol_forward_test.go index 99e9225117..b9e62c3f9d 100644 --- a/backend/internal/service/openai_ws_protocol_forward_test.go +++ b/backend/internal/service/openai_ws_protocol_forward_test.go @@ -619,6 +619,7 @@ func TestNewOpenAIGatewayService_InitializesOpenAIWSResolver(t *testing.T) { nil, nil, nil, + nil, nil, // userPlatformQuotaRepo ) diff --git a/backend/internal/service/scheduler_snapshot_service.go b/backend/internal/service/scheduler_snapshot_service.go index fd48da27af..dc514bc851 100644 --- a/backend/internal/service/scheduler_snapshot_service.go +++ b/backend/internal/service/scheduler_snapshot_service.go @@ -515,7 +515,7 @@ func (s *SchedulerSnapshotService) rebuildByGroupIDs(ctx context.Context, groupI if len(groupIDs) == 0 { return nil } - platforms := []string{PlatformAnthropic, PlatformGemini, PlatformOpenAI, PlatformAntigravity} + platforms := []string{PlatformAnthropic, PlatformGemini, PlatformOpenAI, PlatformAntigravity, PlatformGrok} var firstErr error for _, platform := range platforms { if err := s.rebuildBucketsForPlatform(ctx, platform, groupIDs, reason, seen); err != nil && firstErr == nil { @@ -817,7 +817,7 @@ func (s *SchedulerSnapshotService) fullRebuildInterval() time.Duration { func (s *SchedulerSnapshotService) defaultBuckets(ctx context.Context) ([]SchedulerBucket, error) { buckets := make([]SchedulerBucket, 0) - platforms := []string{PlatformAnthropic, PlatformGemini, PlatformOpenAI, PlatformAntigravity} + platforms := []string{PlatformAnthropic, PlatformGemini, PlatformOpenAI, PlatformAntigravity, PlatformGrok} for _, platform := range platforms { buckets = append(buckets, SchedulerBucket{GroupID: 0, Platform: platform, Mode: SchedulerModeSingle}) buckets = append(buckets, SchedulerBucket{GroupID: 0, Platform: platform, Mode: SchedulerModeForced}) diff --git a/backend/internal/service/token_cache_invalidator.go b/backend/internal/service/token_cache_invalidator.go index 74c9edc399..cf749b7773 100644 --- a/backend/internal/service/token_cache_invalidator.go +++ b/backend/internal/service/token_cache_invalidator.go @@ -44,6 +44,9 @@ func (c *CompositeTokenCacheInvalidator) InvalidateToken(ctx context.Context, ac keysToDelete = append(keysToDelete, "ag:"+accountIDKey) case PlatformOpenAI: keysToDelete = append(keysToDelete, OpenAITokenCacheKey(account)) + case PlatformGrok: + keysToDelete = append(keysToDelete, GrokTokenCacheKey(account)) + keysToDelete = append(keysToDelete, "grok:"+accountIDKey) case PlatformAnthropic: keysToDelete = append(keysToDelete, ClaudeTokenCacheKey(account)) default: diff --git a/backend/internal/service/token_cache_key_test.go b/backend/internal/service/token_cache_key_test.go index 6215eeaf4e..309abe18e9 100644 --- a/backend/internal/service/token_cache_key_test.go +++ b/backend/internal/service/token_cache_key_test.go @@ -199,6 +199,68 @@ func TestOpenAITokenCacheKey(t *testing.T) { } } +func TestGrokTokenCacheKey(t *testing.T) { + tests := []struct { + name string + account *Account + expected string + }{ + { + name: "basic_account", + account: &Account{ + ID: 350, + }, + expected: "grok:account:350", + }, + { + name: "account_with_email_uses_account_id", + account: &Account{ + ID: 351, + Credentials: map[string]any{ + "email": "same-user@example.com", + }, + }, + expected: "grok:account:351", + }, + { + name: "account_id_zero", + account: &Account{ + ID: 0, + }, + expected: "grok:account:0", + }, + { + name: "nil_account", + account: nil, + expected: "grok:account:0", + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + result := GrokTokenCacheKey(tt.account) + require.Equal(t, tt.expected, result) + }) + } +} + +func TestGrokTokenCacheKeySeparatesAccountsWithSameEmail(t *testing.T) { + first := &Account{ + ID: 351, + Credentials: map[string]any{ + "email": "same-user@example.com", + }, + } + second := &Account{ + ID: 352, + Credentials: map[string]any{ + "email": "same-user@example.com", + }, + } + + require.NotEqual(t, GrokTokenCacheKey(first), GrokTokenCacheKey(second)) +} + func TestClaudeTokenCacheKey(t *testing.T) { tests := []struct { name string diff --git a/backend/internal/service/token_refresh_service.go b/backend/internal/service/token_refresh_service.go index 96f5f010d1..08761f8220 100644 --- a/backend/internal/service/token_refresh_service.go +++ b/backend/internal/service/token_refresh_service.go @@ -10,6 +10,7 @@ import ( "time" "github.com/Wei-Shaw/sub2api/internal/config" + "github.com/Wei-Shaw/sub2api/internal/util/logredact" ) // tokenRefreshTempUnschedDuration token 刷新重试耗尽后临时不可调度的持续时间 @@ -49,6 +50,7 @@ func NewTokenRefreshService( schedulerCache SchedulerCache, cfg *config.Config, tempUnschedCache TempUnschedCache, + grokOAuthServices ...*GrokOAuthService, ) *TokenRefreshService { s := &TokenRefreshService{ accountRepo: accountRepo, @@ -65,6 +67,11 @@ func NewTokenRefreshService( claudeRefresher := NewClaudeTokenRefresher(oauthService) geminiRefresher := NewGeminiTokenRefresher(geminiOAuthService) agRefresher := NewAntigravityTokenRefresher(antigravityOAuthService) + var grokOAuthService *GrokOAuthService + if len(grokOAuthServices) > 0 { + grokOAuthService = grokOAuthServices[0] + } + grokRefresher := NewGrokTokenRefresher(grokOAuthService) // 注册平台特定的刷新器(TokenRefresher 接口) s.refreshers = []TokenRefresher{ @@ -72,6 +79,7 @@ func NewTokenRefreshService( openAIRefresher, geminiRefresher, agRefresher, + grokRefresher, } // 注册对应的 OAuthRefreshExecutor(带 CacheKey 方法) @@ -80,6 +88,7 @@ func NewTokenRefreshService( openAIRefresher, geminiRefresher, agRefresher, + grokRefresher, } return s @@ -301,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", @@ -338,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", @@ -442,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) { diff --git a/backend/internal/service/token_refresh_service_test.go b/backend/internal/service/token_refresh_service_test.go index 24adcfb6b3..df14edaf28 100644 --- a/backend/internal/service/token_refresh_service_test.go +++ b/backend/internal/service/token_refresh_service_test.go @@ -20,6 +20,8 @@ type tokenRefreshAccountRepo struct { setErrorCalls int clearTempCalls int setTempUnschedCalls int + lastErrorMessage string + lastTempUnschedReason string lastAccount *Account updateErr error } @@ -51,6 +53,7 @@ func (r *tokenRefreshAccountRepo) UpdateCredentials(ctx context.Context, id int6 func (r *tokenRefreshAccountRepo) SetError(ctx context.Context, id int64, errorMsg string) error { r.setErrorCalls++ + r.lastErrorMessage = errorMsg return nil } @@ -61,6 +64,7 @@ func (r *tokenRefreshAccountRepo) ClearTempUnschedulable(ctx context.Context, id func (r *tokenRefreshAccountRepo) SetTempUnschedulable(ctx context.Context, id int64, until time.Time, reason string) error { r.setTempUnschedCalls++ + r.lastTempUnschedReason = reason return nil } @@ -76,9 +80,13 @@ func (s *tokenCacheInvalidatorStub) InvalidateToken(ctx context.Context, account type tempUnschedCacheStub struct { deleteCalls int + setCalls int + lastState *TempUnschedState } func (s *tempUnschedCacheStub) SetTempUnsched(ctx context.Context, accountID int64, state *TempUnschedState) error { + s.setCalls++ + s.lastState = state return nil } @@ -538,6 +546,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}, } diff --git a/backend/internal/service/wire.go b/backend/internal/service/wire.go index bc0cf46f35..b3b1170d97 100644 --- a/backend/internal/service/wire.go +++ b/backend/internal/service/wire.go @@ -63,6 +63,7 @@ func ProvideTokenRefreshService( openaiOAuthService *OpenAIOAuthService, geminiOAuthService *GeminiOAuthService, antigravityOAuthService *AntigravityOAuthService, + grokOAuthService *GrokOAuthService, cacheInvalidator TokenCacheInvalidator, schedulerCache SchedulerCache, cfg *config.Config, @@ -72,7 +73,7 @@ func ProvideTokenRefreshService( refreshAPI *OAuthRefreshAPI, runtimeBlocker AccountRuntimeBlocker, ) *TokenRefreshService { - svc := NewTokenRefreshService(accountRepo, oauthService, openaiOAuthService, geminiOAuthService, antigravityOAuthService, cacheInvalidator, schedulerCache, cfg, tempUnschedCache) + svc := NewTokenRefreshService(accountRepo, oauthService, openaiOAuthService, geminiOAuthService, antigravityOAuthService, cacheInvalidator, schedulerCache, cfg, tempUnschedCache, grokOAuthService) // 注入 OpenAI privacy opt-out 依赖 svc.SetPrivacyDeps(privacyClientFactory, proxyRepo) // 注入统一 OAuth 刷新 API(消除 TokenRefreshService 与 TokenProvider 之间的竞争条件) @@ -124,6 +125,15 @@ func ProvideOpenAIQuotaService( return NewOpenAIQuotaService(accountRepo, proxyRepo, tokenProvider, privacyClientFactory) } +func ProvideGrokQuotaService( + accountRepo AccountRepository, + proxyRepo ProxyRepository, + tokenProvider *GrokTokenProvider, + httpUpstream HTTPUpstream, +) *GrokQuotaService { + return NewGrokQuotaService(accountRepo, proxyRepo, tokenProvider, httpUpstream) +} + // ProvideGeminiTokenProvider creates GeminiTokenProvider with OAuthRefreshAPI injection func ProvideGeminiTokenProvider( accountRepo AccountRepository, @@ -154,6 +164,22 @@ func ProvideAntigravityTokenProvider( return p } +// ProvideGrokTokenProvider creates GrokTokenProvider with OAuthRefreshAPI injection. +func ProvideGrokTokenProvider( + accountRepo AccountRepository, + tokenCache GeminiTokenCache, + grokOAuthService *GrokOAuthService, + refreshAPI *OAuthRefreshAPI, + tempUnschedCache TempUnschedCache, +) *GrokTokenProvider { + p := NewGrokTokenProvider(accountRepo, tokenCache) + executor := NewGrokTokenRefresher(grokOAuthService) + p.SetRefreshAPI(refreshAPI, executor) + p.SetRefreshPolicy(AntigravityProviderRefreshPolicy()) + p.SetTempUnschedCache(tempUnschedCache) + return p +} + // ProvideDashboardAggregationService 创建并启动仪表盘聚合服务 func ProvideDashboardAggregationService(repo DashboardAggregationRepository, timingWheel *TimingWheelService, lockCache LeaderLockCache, db *sql.DB, cfg *config.Config) *DashboardAggregationService { svc := NewDashboardAggregationService(repo, timingWheel, cfg) @@ -535,6 +561,7 @@ var ProviderSet = wire.NewSet( wire.Bind(new(AccountRuntimeBlocker), new(*OpenAIGatewayService)), NewOAuthService, ProvideOpenAIOAuthService, + NewGrokOAuthService, NewGeminiOAuthService, NewGeminiQuotaService, NewCompositeTokenCacheInvalidator, @@ -544,8 +571,10 @@ var ProviderSet = wire.NewSet( ProvideGeminiTokenProvider, NewGeminiMessagesCompatService, ProvideAntigravityTokenProvider, + ProvideGrokTokenProvider, ProvideOpenAITokenProvider, ProvideOpenAIQuotaService, + ProvideGrokQuotaService, ProvideClaudeTokenProvider, NewAntigravityGatewayService, ProvideRateLimitService, @@ -583,6 +612,7 @@ var ProviderSet = wire.NewSet( ProvideUsageCleanupService, ProvideDeferredService, NewAntigravityQuotaFetcher, + NewGrokQuotaFetcher, NewUserAttributeService, NewUsageCache, NewTotpService, diff --git a/frontend/src/api/__tests__/settings.authSourceDefaults.spec.ts b/frontend/src/api/__tests__/settings.authSourceDefaults.spec.ts index 4dc7db2ef9..deeb758250 100644 --- a/frontend/src/api/__tests__/settings.authSourceDefaults.spec.ts +++ b/frontend/src/api/__tests__/settings.authSourceDefaults.spec.ts @@ -9,12 +9,13 @@ import { type DefaultPlatformQuotasMap, } from "@/api/admin/settings"; -/** 全 null 的 4 平台 map,用于断言归一化默认值 */ +/** 全 null 的 5 平台 map,用于断言归一化默认值 */ const allNullQuotas: DefaultPlatformQuotasMap = { anthropic: { daily: null, weekly: null, monthly: null }, openai: { daily: null, weekly: null, monthly: null }, gemini: { daily: null, weekly: null, monthly: null }, antigravity: { daily: null, weekly: null, monthly: null }, + grok: { daily: null, weekly: null, monthly: null }, } describe("admin settings auth source defaults helpers", () => { @@ -236,11 +237,12 @@ describe("normalizePlatformQuotasMap", () => { expect(result.openai).toEqual({ daily: null, weekly: null, monthly: null }); expect(result.gemini).toEqual({ daily: null, weekly: null, monthly: null }); expect(result.antigravity).toEqual({ daily: null, weekly: null, monthly: null }); + expect(result.grok).toEqual({ daily: null, weekly: null, monthly: null }); }); - it("无参数时返回全 4 平台全 null", () => { + it("无参数时返回全 5 平台全 null", () => { const result = normalizePlatformQuotasMap(); - expect(Object.keys(result)).toHaveLength(4); + expect(Object.keys(result)).toHaveLength(5); for (const v of Object.values(result)) { expect(v).toEqual({ daily: null, weekly: null, monthly: null }); } @@ -288,7 +290,7 @@ describe("sanitizePlatformQuotasMap", () => { it("缺失平台填充为全 null", () => { const result = sanitizePlatformQuotasMap({}); - expect(Object.keys(result)).toHaveLength(4); + expect(Object.keys(result)).toHaveLength(5); for (const v of Object.values(result)) { expect(v).toEqual({ daily: null, weekly: null, monthly: null }); } diff --git a/frontend/src/api/admin/grok.ts b/frontend/src/api/admin/grok.ts new file mode 100644 index 0000000000..2e0f5349b0 --- /dev/null +++ b/frontend/src/api/admin/grok.ts @@ -0,0 +1,121 @@ +/** + * Admin Grok/xAI API endpoints + * Handles xAI OAuth flows for administrators. + */ + +import { apiClient } from '../client' + +export interface GrokAuthUrlResponse { + auth_url: string + session_id: string + state: string +} + +export interface GrokAuthUrlRequest { + proxy_id?: number + redirect_uri?: string +} + +export interface GrokExchangeCodeRequest { + session_id: string + state: string + code: string + proxy_id?: number + redirect_uri?: string +} + +export interface GrokTokenInfo { + access_token?: string + refresh_token?: string + token_type?: string + id_token?: string + expires_at?: number | string + expires_in?: number + scope?: string + client_id?: string + email?: string + subscription_tier?: string + entitlement_status?: string + [key: string]: unknown +} + +export interface GrokQuotaWindow { + limit?: number | null + remaining?: number | null + reset_unix?: number | null + reset_at?: string | null +} + +export interface GrokQuotaSnapshot { + requests?: GrokQuotaWindow | null + tokens?: GrokQuotaWindow | null + retry_after_seconds?: number | null + subscription_tier?: string + entitlement_status?: string + status_code?: number + headers?: Record + headers_observed: boolean + observation_source?: string + last_probe_at?: string + last_headers_seen_at?: string + updated_at: string +} + +export interface GrokQuotaProbeResult { + source: 'active_probe' + snapshot?: GrokQuotaSnapshot | null + status_code?: number + headers_observed: boolean + reset_supported: boolean + fetched_at: number +} + +export interface GrokQuotaResetResult { + supported: boolean + code: string + message: string +} + +export async function generateAuthUrl( + payload: GrokAuthUrlRequest +): Promise { + const { data } = await apiClient.post( + '/admin/grok/oauth/auth-url', + payload + ) + return data +} + +export async function exchangeCode(payload: GrokExchangeCodeRequest): Promise { + const { data } = await apiClient.post( + '/admin/grok/oauth/exchange-code', + payload + ) + return data +} + +export async function refreshGrokToken( + refreshToken: string, + proxyId?: number | null +): Promise { + const payload: Record = { refresh_token: refreshToken } + if (proxyId) payload.proxy_id = proxyId + + const { data } = await apiClient.post( + '/admin/grok/oauth/refresh-token', + payload + ) + return data +} + +export async function queryQuota(id: number): Promise { + const { data } = await apiClient.get(`/admin/grok/accounts/${id}/quota`) + return data +} + +export async function resetQuota(id: number): Promise { + const { data } = await apiClient.post(`/admin/grok/accounts/${id}/reset-quota`) + return data +} + +export default { generateAuthUrl, exchangeCode, refreshGrokToken, queryQuota, resetQuota } diff --git a/frontend/src/api/admin/index.ts b/frontend/src/api/admin/index.ts index 176498a287..1698a92c43 100644 --- a/frontend/src/api/admin/index.ts +++ b/frontend/src/api/admin/index.ts @@ -17,6 +17,7 @@ import subscriptionsAPI from './subscriptions' import usageAPI from './usage' import geminiAPI from './gemini' import antigravityAPI from './antigravity' +import grokAPI from './grok' import userAttributesAPI from './userAttributes' import opsAPI from './ops' import errorPassthroughAPI from './errorPassthrough' @@ -51,6 +52,7 @@ export const adminAPI = { usage: usageAPI, gemini: geminiAPI, antigravity: antigravityAPI, + grok: grokAPI, userAttributes: userAttributesAPI, ops: opsAPI, errorPassthrough: errorPassthroughAPI, @@ -83,6 +85,7 @@ export { usageAPI, geminiAPI, antigravityAPI, + grokAPI, userAttributesAPI, opsAPI, errorPassthroughAPI, diff --git a/frontend/src/api/admin/settings.ts b/frontend/src/api/admin/settings.ts index 4c1c4685c4..bc3f3fa4df 100644 --- a/frontend/src/api/admin/settings.ts +++ b/frontend/src/api/admin/settings.ts @@ -17,7 +17,7 @@ export interface DefaultSubscriptionSetting { } // ── 平台限额类型 ────────────────────────────────────────────────── -export type PlatformType = "anthropic" | "openai" | "gemini" | "antigravity" +export type PlatformType = "anthropic" | "openai" | "gemini" | "antigravity" | "grok" export type QuotaWindowType = "daily" | "weekly" | "monthly" /** 单平台三档限额;null = 不限制,undefined = 未填(等价 null) */ @@ -30,7 +30,7 @@ export interface PlatformQuotaLimits { /** 全平台默认限额 map(key = PlatformType) */ export type DefaultPlatformQuotasMap = Partial> -const PLATFORMS: PlatformType[] = ["anthropic", "openai", "gemini", "antigravity"] +const PLATFORMS: PlatformType[] = ["anthropic", "openai", "gemini", "antigravity", "grok"] /** 归一化为全 4 平台 × 3 窗口(缺失填 null),供模板非空绑定 */ export function normalizePlatformQuotasMap(input?: DefaultPlatformQuotasMap | null): DefaultPlatformQuotasMap { diff --git a/frontend/src/api/admin/users.ts b/frontend/src/api/admin/users.ts index 7667dfc4c9..8ff022ff68 100644 --- a/frontend/src/api/admin/users.ts +++ b/frontend/src/api/admin/users.ts @@ -307,7 +307,7 @@ export async function bindUserAuthIdentity( /** * Platform quota types */ -export type PlatformQuotaPlatform = 'anthropic' | 'openai' | 'gemini' | 'antigravity' +export type PlatformQuotaPlatform = 'anthropic' | 'openai' | 'gemini' | 'antigravity' | 'grok' export type PlatformQuotaWindow = 'daily' | 'weekly' | 'monthly' export interface PlatformQuotaItem { diff --git a/frontend/src/components/account/AccountUsageCell.vue b/frontend/src/components/account/AccountUsageCell.vue index 156b4e2497..b893c1e801 100644 --- a/frontend/src/components/account/AccountUsageCell.vue +++ b/frontend/src/components/account/AccountUsageCell.vue @@ -102,7 +102,11 @@ -
-
+
+
-
+ + +
@@ -320,6 +324,78 @@
-
+ + +