Merge pull request #3310 from heathermhuang/codex/grok-subscription-support

feat: add grok subscription support
This commit is contained in:
Wesley Liddick
2026-06-26 15:41:52 +08:00
committed by GitHub
109 changed files with 5594 additions and 224 deletions
+7
View File
@@ -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()
+18 -5
View File
@@ -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()
+1
View File
@@ -74,6 +74,7 @@ func TestProvideCleanup_WithMinimalDependencies_NoPanic(t *testing.T) {
openAIOAuthSvc,
geminiOAuthSvc,
antigravityOAuthSvc,
nil, // grokOAuth
nil, // openAIGateway
nil, // scheduledTestRunner
nil, // backupSvc
+1 -1
View File
@@ -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)
+2
View File
@@ -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=
+1
View File
@@ -22,6 +22,7 @@ const (
PlatformOpenAI = "openai"
PlatformGemini = "gemini"
PlatformAntigravity = "antigravity"
PlatformGrok = "grok"
)
// Account type constants
@@ -509,6 +509,7 @@ var platformToLiteLLMProvider = map[string]string{
service.PlatformOpenAI: "openai",
service.PlatformGemini: "google",
service.PlatformAntigravity: "anthropic",
service.PlatformGrok: "xai",
}
// SyncPricingModels 返回 LiteLLM 定价目录中指定平台的最新模型列表
@@ -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())
}
@@ -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")
}
@@ -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"`
@@ -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)
}
}
+1 -1
View File
@@ -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
}
@@ -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 {
+1
View File
@@ -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
@@ -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",
@@ -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",
@@ -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,
@@ -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)
+3
View File
@@ -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,
@@ -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 验证规则配置的有效性
+46
View File
@@ -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
}
+448
View File
@@ -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"`
}
+195
View File
@@ -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"])
}
+181
View File
@@ -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 ""
}
+63
View File
@@ -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)
}
@@ -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)
}
@@ -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")
}
@@ -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 {
+1
View File
@@ -141,6 +141,7 @@ var ProviderSet = wire.NewSet(
NewClaudeOAuthClient,
NewHTTPUpstream,
NewOpenAIOAuthClient,
NewGrokOAuthClient,
NewGeminiOAuthClient,
NewGeminiCliCodeAssistClient,
NewGeminiDriveClient,
+17
View File
@@ -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")
{
+61 -10
View File
@@ -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
}
+44 -2
View File
@@ -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)
}
}
+54 -1
View File
@@ -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
}
+5 -2
View File
@@ -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)
}
@@ -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) {
@@ -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))
}
+19 -4
View File
@@ -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 {
@@ -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
}
@@ -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 标识。
@@ -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)
}
@@ -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)
}
@@ -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
}
}
@@ -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)
})
}
}
@@ -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
}
@@ -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)
})
}
}
@@ -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)
}
@@ -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")
}
@@ -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
}
+13
View File
@@ -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)
@@ -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
}
@@ -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))
}
@@ -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,
@@ -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()
@@ -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) {
@@ -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 协议)。
//
@@ -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
}
@@ -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)
}
@@ -224,6 +224,7 @@ func newOpenAIRecordUsageServiceForTest(usageRepo UsageLogRepository, userRepo U
nil,
nil,
nil,
nil,
nil, // userPlatformQuotaRepo
)
svc.userGroupRateResolver = newUserGroupRateResolver(
@@ -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)
}
@@ -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{
@@ -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
@@ -619,6 +619,7 @@ func TestNewOpenAIGatewayService_InitializesOpenAIWSResolver(t *testing.T) {
nil,
nil,
nil,
nil,
nil, // userPlatformQuotaRepo
)
@@ -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})
@@ -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:
@@ -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
@@ -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) {
@@ -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},
}
+31 -1
View File
@@ -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,