mirror of
https://github.com/Wei-Shaw/sub2api.git
synced 2026-09-24 16:05:44 +08:00
Merge pull request #3310 from heathermhuang/codex/grok-subscription-support
feat: add grok subscription support
This commit is contained in:
@@ -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()
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -74,6 +74,7 @@ func TestProvideCleanup_WithMinimalDependencies_NoPanic(t *testing.T) {
|
||||
openAIOAuthSvc,
|
||||
geminiOAuthSvc,
|
||||
antigravityOAuthSvc,
|
||||
nil, // grokOAuth
|
||||
nil, // openAIGateway
|
||||
nil, // scheduledTestRunner
|
||||
nil, // backupSvc
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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=
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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 验证规则配置的有效性
|
||||
|
||||
@@ -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
|
||||
}
|
||||
@@ -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"`
|
||||
}
|
||||
@@ -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"])
|
||||
}
|
||||
@@ -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 ""
|
||||
}
|
||||
@@ -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 {
|
||||
|
||||
@@ -141,6 +141,7 @@ var ProviderSet = wire.NewSet(
|
||||
NewClaudeOAuthClient,
|
||||
NewHTTPUpstream,
|
||||
NewOpenAIOAuthClient,
|
||||
NewGrokOAuthClient,
|
||||
NewGeminiOAuthClient,
|
||||
NewGeminiCliCodeAssistClient,
|
||||
NewGeminiDriveClient,
|
||||
|
||||
@@ -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")
|
||||
{
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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))
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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},
|
||||
}
|
||||
|
||||
@@ -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,
|
||||
|
||||
Reference in New Issue
Block a user