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

feat: add grok subscription support
This commit is contained in:
Wesley Liddick
2026-06-26 15:41:52 +08:00
committed by GitHub
109 changed files with 5594 additions and 224 deletions
@@ -509,6 +509,7 @@ var platformToLiteLLMProvider = map[string]string{
service.PlatformOpenAI: "openai",
service.PlatformGemini: "google",
service.PlatformAntigravity: "anthropic",
service.PlatformGrok: "xai",
}
// SyncPricingModels 返回 LiteLLM 定价目录中指定平台的最新模型列表
@@ -0,0 +1,246 @@
package admin
import (
"strconv"
"strings"
"github.com/Wei-Shaw/sub2api/internal/handler/dto"
"github.com/Wei-Shaw/sub2api/internal/pkg/response"
"github.com/Wei-Shaw/sub2api/internal/pkg/xai"
"github.com/Wei-Shaw/sub2api/internal/service"
"github.com/gin-gonic/gin"
)
type GrokOAuthHandler struct {
grokOAuthService *service.GrokOAuthService
adminService service.AdminService
quotaService *service.GrokQuotaService
}
func NewGrokOAuthHandler(
grokOAuthService *service.GrokOAuthService,
adminService service.AdminService,
quotaService *service.GrokQuotaService,
) *GrokOAuthHandler {
return &GrokOAuthHandler{
grokOAuthService: grokOAuthService,
adminService: adminService,
quotaService: quotaService,
}
}
type GrokGenerateAuthURLRequest struct {
ProxyID *int64 `json:"proxy_id"`
RedirectURI string `json:"redirect_uri"`
}
func (h *GrokOAuthHandler) GenerateAuthURL(c *gin.Context) {
var req GrokGenerateAuthURLRequest
if err := c.ShouldBindJSON(&req); err != nil {
req = GrokGenerateAuthURLRequest{}
}
result, err := h.grokOAuthService.GenerateAuthURL(c.Request.Context(), req.ProxyID, req.RedirectURI)
if err != nil {
response.ErrorFrom(c, err)
return
}
response.Success(c, result)
}
type GrokExchangeCodeRequest struct {
SessionID string `json:"session_id" binding:"required"`
Code string `json:"code" binding:"required"`
State string `json:"state"`
RedirectURI string `json:"redirect_uri"`
ProxyID *int64 `json:"proxy_id"`
}
func (h *GrokOAuthHandler) ExchangeCode(c *gin.Context) {
var req GrokExchangeCodeRequest
if err := c.ShouldBindJSON(&req); err != nil {
response.BadRequest(c, "Invalid request: "+err.Error())
return
}
tokenInfo, err := h.grokOAuthService.ExchangeCode(c.Request.Context(), &service.GrokExchangeCodeInput{
SessionID: req.SessionID,
Code: req.Code,
State: req.State,
RedirectURI: req.RedirectURI,
ProxyID: req.ProxyID,
})
if err != nil {
response.ErrorFrom(c, err)
return
}
response.Success(c, tokenInfo)
}
type GrokRefreshTokenRequest struct {
RefreshToken string `json:"refresh_token"`
RT string `json:"rt"`
ClientID string `json:"client_id"`
ProxyID *int64 `json:"proxy_id"`
}
func (h *GrokOAuthHandler) RefreshToken(c *gin.Context) {
var req GrokRefreshTokenRequest
if err := c.ShouldBindJSON(&req); err != nil {
response.BadRequest(c, "Invalid request: "+err.Error())
return
}
refreshToken := strings.TrimSpace(req.RefreshToken)
if refreshToken == "" {
refreshToken = strings.TrimSpace(req.RT)
}
if refreshToken == "" {
response.BadRequest(c, "refresh_token is required")
return
}
var proxyURL string
if req.ProxyID != nil {
proxy, err := h.adminService.GetProxy(c.Request.Context(), *req.ProxyID)
if err == nil && proxy != nil {
proxyURL = proxy.URL()
}
}
tokenInfo, err := h.grokOAuthService.RefreshToken(c.Request.Context(), refreshToken, proxyURL, req.ClientID)
if err != nil {
response.ErrorFrom(c, err)
return
}
response.Success(c, tokenInfo)
}
func (h *GrokOAuthHandler) RefreshAccountToken(c *gin.Context) {
accountID, err := strconv.ParseInt(c.Param("id"), 10, 64)
if err != nil {
response.BadRequest(c, "Invalid account ID")
return
}
account, err := h.adminService.GetAccount(c.Request.Context(), accountID)
if err != nil {
response.ErrorFrom(c, err)
return
}
if account.Platform != service.PlatformGrok {
response.BadRequest(c, "Account platform does not match Grok OAuth endpoint")
return
}
if !account.IsOAuth() {
response.BadRequest(c, "Cannot refresh non-OAuth account credentials")
return
}
tokenInfo, err := h.grokOAuthService.RefreshAccountToken(c.Request.Context(), account)
if err != nil {
response.ErrorFrom(c, err)
return
}
newCredentials := h.grokOAuthService.BuildAccountCredentials(tokenInfo)
newCredentials = service.MergeCredentials(account.Credentials, newCredentials)
if baseURL := strings.TrimSpace(account.GetCredential("base_url")); baseURL != "" {
newCredentials["base_url"] = baseURL
}
updatedAccount, err := h.adminService.UpdateAccount(c.Request.Context(), accountID, &service.UpdateAccountInput{
Credentials: newCredentials,
})
if err != nil {
response.ErrorFrom(c, err)
return
}
response.Success(c, dto.AccountFromService(updatedAccount))
}
func (h *GrokOAuthHandler) CreateAccountFromOAuth(c *gin.Context) {
var req struct {
SessionID string `json:"session_id" binding:"required"`
Code string `json:"code" binding:"required"`
State string `json:"state"`
RedirectURI string `json:"redirect_uri"`
ProxyID *int64 `json:"proxy_id"`
Name string `json:"name"`
Concurrency int `json:"concurrency"`
Priority int `json:"priority"`
GroupIDs []int64 `json:"group_ids"`
}
if err := c.ShouldBindJSON(&req); err != nil {
response.BadRequest(c, "Invalid request: "+err.Error())
return
}
tokenInfo, err := h.grokOAuthService.ExchangeCode(c.Request.Context(), &service.GrokExchangeCodeInput{
SessionID: req.SessionID,
Code: req.Code,
State: req.State,
RedirectURI: req.RedirectURI,
ProxyID: req.ProxyID,
})
if err != nil {
response.ErrorFrom(c, err)
return
}
credentials := h.grokOAuthService.BuildAccountCredentials(tokenInfo)
name := strings.TrimSpace(req.Name)
if name == "" && tokenInfo.Email != "" {
name = tokenInfo.Email
}
if name == "" {
name = "Grok OAuth Account"
}
account, err := h.adminService.CreateAccount(c.Request.Context(), &service.CreateAccountInput{
Name: name,
Platform: service.PlatformGrok,
Type: service.AccountTypeOAuth,
Credentials: credentials,
ProxyID: req.ProxyID,
Concurrency: req.Concurrency,
Priority: req.Priority,
GroupIDs: req.GroupIDs,
})
if err != nil {
response.ErrorFrom(c, err)
return
}
response.Success(c, dto.AccountFromService(account))
}
func (h *GrokOAuthHandler) QueryQuota(c *gin.Context) {
accountID, err := strconv.ParseInt(c.Param("id"), 10, 64)
if err != nil {
response.BadRequest(c, "Invalid account ID")
return
}
if h.quotaService == nil {
response.BadRequest(c, "grok quota service is not enabled")
return
}
result, err := h.quotaService.ProbeUsage(c.Request.Context(), accountID)
if err != nil {
response.ErrorFrom(c, err)
return
}
response.Success(c, result)
}
func (h *GrokOAuthHandler) ResetQuota(c *gin.Context) {
accountID, err := strconv.ParseInt(c.Param("id"), 10, 64)
if err != nil {
response.BadRequest(c, "Invalid account ID")
return
}
if h.quotaService == nil {
response.BadRequest(c, "grok quota service is not enabled")
return
}
result, err := h.quotaService.ResetQuota(c.Request.Context(), accountID)
if err != nil {
response.ErrorFrom(c, err)
return
}
response.Success(c, result)
}
func (h *GrokOAuthHandler) RuntimeSanity(c *gin.Context) {
response.Success(c, xai.RuntimeSanity())
}
@@ -0,0 +1,147 @@
//go:build unit
package admin
import (
"context"
"io"
"net/http"
"net/http/httptest"
"strings"
"testing"
"time"
"github.com/gin-gonic/gin"
"github.com/stretchr/testify/require"
"github.com/Wei-Shaw/sub2api/internal/pkg/tlsfingerprint"
"github.com/Wei-Shaw/sub2api/internal/pkg/xai"
"github.com/Wei-Shaw/sub2api/internal/service"
)
type grokQuotaHandlerAccountRepo struct {
service.AccountRepository
account *service.Account
updates map[int64]map[string]any
}
func (r *grokQuotaHandlerAccountRepo) GetByID(_ context.Context, id int64) (*service.Account, error) {
if r.account != nil && r.account.ID == id {
return r.account, nil
}
return nil, service.ErrAccountNotFound
}
func (r *grokQuotaHandlerAccountRepo) UpdateExtra(_ context.Context, id int64, updates map[string]any) error {
if r.updates == nil {
r.updates = make(map[int64]map[string]any)
}
r.updates[id] = updates
return nil
}
type grokQuotaHandlerUpstream struct {
resp *http.Response
lastReq *http.Request
lastBody []byte
}
func (u *grokQuotaHandlerUpstream) Do(req *http.Request, _ string, _ int64, _ int) (*http.Response, error) {
u.lastReq = req
if req.Body != nil {
u.lastBody, _ = io.ReadAll(req.Body)
}
return u.resp, nil
}
func (u *grokQuotaHandlerUpstream) DoWithTLS(
req *http.Request,
proxyURL string,
accountID int64,
accountConcurrency int,
_ *tlsfingerprint.Profile,
) (*http.Response, error) {
return u.Do(req, proxyURL, accountID, accountConcurrency)
}
func TestGrokOAuthHandlerQueryQuotaProbesUpstream(t *testing.T) {
gin.SetMode(gin.TestMode)
repo := &grokQuotaHandlerAccountRepo{account: &service.Account{
ID: 42,
Platform: service.PlatformGrok,
Type: service.AccountTypeOAuth,
Concurrency: 1,
Credentials: map[string]any{
"access_token": "access-token",
"expires_at": time.Now().Add(time.Hour).UTC().Format(time.RFC3339),
},
}}
upstream := &grokQuotaHandlerUpstream{resp: &http.Response{
StatusCode: http.StatusOK,
Header: http.Header{
"X-Ratelimit-Limit-Requests": []string{"10"},
"X-Ratelimit-Remaining-Requests": []string{"8"},
},
Body: io.NopCloser(strings.NewReader(`{"id":"resp_probe"}`)),
}}
quotaService := service.NewGrokQuotaService(repo, nil, service.NewGrokTokenProvider(repo, nil), upstream)
handler := NewGrokOAuthHandler(nil, nil, quotaService)
router := gin.New()
router.GET("/api/v1/admin/grok/accounts/:id/quota", handler.QueryQuota)
rec := httptest.NewRecorder()
req := httptest.NewRequest(http.MethodGet, "/api/v1/admin/grok/accounts/42/quota", nil)
router.ServeHTTP(rec, req)
require.Equal(t, http.StatusOK, rec.Code)
require.Contains(t, rec.Body.String(), `"source":"active_probe"`)
require.Contains(t, rec.Body.String(), `"headers_observed":true`)
require.NotContains(t, rec.Body.String(), "access-token")
require.Equal(t, xai.DefaultBaseURL+"/responses", upstream.lastReq.URL.String())
require.Equal(t, "Bearer access-token", upstream.lastReq.Header.Get("Authorization"))
require.Contains(t, string(upstream.lastBody), `"store":false`)
require.NotNil(t, repo.updates[42])
}
func TestGrokOAuthHandlerResetQuotaReturnsUnsupported(t *testing.T) {
gin.SetMode(gin.TestMode)
repo := &grokQuotaHandlerAccountRepo{account: &service.Account{
ID: 43,
Platform: service.PlatformGrok,
Type: service.AccountTypeOAuth,
}}
quotaService := service.NewGrokQuotaService(repo, nil, nil, nil)
handler := NewGrokOAuthHandler(nil, nil, quotaService)
router := gin.New()
router.POST("/api/v1/admin/grok/accounts/:id/reset-quota", handler.ResetQuota)
rec := httptest.NewRecorder()
req := httptest.NewRequest(http.MethodPost, "/api/v1/admin/grok/accounts/43/reset-quota", nil)
router.ServeHTTP(rec, req)
require.Equal(t, http.StatusNotImplemented, rec.Code)
require.Contains(t, rec.Body.String(), `"reason":"GROK_QUOTA_RESET_UNSUPPORTED"`)
require.NotContains(t, rec.Body.String(), "access-token")
}
func TestGrokOAuthHandlerRuntimeSanityDoesNotExposeSecrets(t *testing.T) {
gin.SetMode(gin.TestMode)
t.Setenv(xai.EnvBaseURL, "http://127.0.0.1:8080/v1?access_token=secret")
t.Setenv(xai.EnvClientID, "client-secret-like-value")
handler := NewGrokOAuthHandler(nil, nil, nil)
router := gin.New()
router.GET("/api/v1/admin/grok/runtime-sanity", handler.RuntimeSanity)
rec := httptest.NewRecorder()
req := httptest.NewRequest(http.MethodGet, "/api/v1/admin/grok/runtime-sanity", nil)
router.ServeHTTP(rec, req)
require.Equal(t, http.StatusOK, rec.Code)
require.Contains(t, rec.Body.String(), `"public_gateway_scope":"responses_only"`)
require.Contains(t, rec.Body.String(), `"valid":false`)
require.NotContains(t, rec.Body.String(), "access_token")
require.NotContains(t, rec.Body.String(), "secret")
require.NotContains(t, rec.Body.String(), "client-secret-like-value")
}
@@ -84,7 +84,7 @@ func NewGroupHandler(adminService service.AdminService, dashboardService *servic
type CreateGroupRequest struct {
Name string `json:"name" binding:"required"`
Description string `json:"description"`
Platform string `json:"platform" binding:"omitempty,oneof=anthropic openai gemini antigravity"`
Platform string `json:"platform" binding:"omitempty,oneof=anthropic openai gemini antigravity grok"`
RateMultiplier float64 `json:"rate_multiplier"`
IsExclusive bool `json:"is_exclusive"`
SubscriptionType string `json:"subscription_type" binding:"omitempty,oneof=standard subscription"`
@@ -124,7 +124,7 @@ type CreateGroupRequest struct {
type UpdateGroupRequest struct {
Name string `json:"name"`
Description *string `json:"description"`
Platform string `json:"platform" binding:"omitempty,oneof=anthropic openai gemini antigravity"`
Platform string `json:"platform" binding:"omitempty,oneof=anthropic openai gemini antigravity grok"`
RateMultiplier *float64 `json:"rate_multiplier"`
IsExclusive *bool `json:"is_exclusive"`
Status string `json:"status" binding:"omitempty,oneof=active inactive"`
@@ -112,9 +112,9 @@ func TestUpdateUserPlatformQuotas_Success(t *testing.T) {
if repo.upsertCalls[0].userID != 42 || len(repo.upsertCalls[0].records) != 2 {
t.Errorf("unexpected upsert call: %+v", repo.upsertCalls[0])
}
// 缓存失效:请求中 2 个 platform + 软删除的 2 个 platform(gemini, antigravity)= 4 次
if len(cache.deleteCalls) != 4 {
t.Errorf("expected 4 cache delete calls, got %d: %+v", len(cache.deleteCalls), cache.deleteCalls)
// 缓存失效:请求中 2 个 platform + 软删除的 3 个 platform(gemini, antigravity, grok)= 5 次
if len(cache.deleteCalls) != 5 {
t.Errorf("expected 5 cache delete calls, got %d: %+v", len(cache.deleteCalls), cache.deleteCalls)
}
}
+1 -1
View File
@@ -77,7 +77,7 @@ func DeriveUpstreamEndpoint(inbound, rawRequestPath, platform string) string {
inbound = strings.TrimSpace(inbound)
switch platform {
case service.PlatformOpenAI:
case service.PlatformOpenAI, service.PlatformGrok:
if inbound == EndpointEmbeddings || inbound == EndpointImagesGenerations || inbound == EndpointImagesEdits {
return inbound
}
@@ -25,6 +25,7 @@ import (
"github.com/Wei-Shaw/sub2api/internal/pkg/logger"
"github.com/Wei-Shaw/sub2api/internal/pkg/openai"
"github.com/Wei-Shaw/sub2api/internal/pkg/timezone"
"github.com/Wei-Shaw/sub2api/internal/pkg/xai"
middleware2 "github.com/Wei-Shaw/sub2api/internal/server/middleware"
"github.com/Wei-Shaw/sub2api/internal/service"
@@ -1157,6 +1158,8 @@ func defaultModelIDsForPlatform(platform string) []string {
ids = append(ids, model.ID)
}
return ids
case service.PlatformGrok:
return xai.DefaultModelIDs()
default:
ids := make([]string, 0, len(claude.DefaultModels))
for _, model := range claude.DefaultModels {
+1
View File
@@ -17,6 +17,7 @@ type AdminHandlers struct {
OpenAIOAuth *admin.OpenAIOAuthHandler
GeminiOAuth *admin.GeminiOAuthHandler
AntigravityOAuth *admin.AntigravityOAuthHandler
GrokOAuth *admin.GrokOAuthHandler
Proxy *admin.ProxyHandler
Redeem *admin.RedeemHandler
Promo *admin.PromoHandler
@@ -101,6 +101,7 @@ func (h *OpenAIGatewayHandler) ChatCompletions(c *gin.Context) {
}
subscription, _ := middleware2.GetSubscriptionFromContext(c)
requestPlatform := openAICompatibleRequestPlatform(apiKey)
service.SetOpsLatencyMs(c, service.OpsAuthLatencyMsKey, time.Since(requestStart).Milliseconds())
routingStart := time.Now()
@@ -144,6 +145,7 @@ func (h *OpenAIGatewayHandler) ChatCompletions(c *gin.Context) {
service.OpenAIUpstreamTransportAny,
service.OpenAIEndpointCapabilityChatCompletions,
false,
requestPlatform,
)
if err != nil {
reqLog.Warn("openai_chat_completions.account_select_failed",
@@ -97,6 +97,13 @@ func wrapUsageRecordTaskContext(parent context.Context, task service.UsageRecord
}
}
func openAICompatibleRequestPlatform(apiKey *service.APIKey) string {
if apiKey != nil && apiKey.Group != nil && apiKey.Group.Platform == service.PlatformGrok {
return service.PlatformGrok
}
return service.PlatformOpenAI
}
// NewOpenAIGatewayHandler creates a new OpenAIGatewayHandler
func NewOpenAIGatewayHandler(
gatewayService *service.OpenAIGatewayService,
@@ -282,6 +289,7 @@ func (h *OpenAIGatewayHandler) Responses(c *gin.Context) {
// Get subscription info (may be nil)
subscription, _ := middleware2.GetSubscriptionFromContext(c)
requestPlatform := openAICompatibleRequestPlatform(apiKey)
service.SetOpsLatencyMs(c, service.OpsAuthLatencyMsKey, time.Since(requestStart).Milliseconds())
routingStart := time.Now()
@@ -332,6 +340,7 @@ func (h *OpenAIGatewayHandler) Responses(c *gin.Context) {
service.OpenAIUpstreamTransportAny,
service.OpenAIEndpointCapabilityChatCompletions,
requireCompact,
requestPlatform,
)
if err != nil {
reqLog.Warn("openai.account_select_failed",
@@ -708,6 +717,7 @@ func (h *OpenAIGatewayHandler) Messages(c *gin.Context) {
}
subscription, _ := middleware2.GetSubscriptionFromContext(c)
requestPlatform := openAICompatibleRequestPlatform(apiKey)
service.SetOpsLatencyMs(c, service.OpsAuthLatencyMsKey, time.Since(requestStart).Milliseconds())
routingStart := time.Now()
@@ -760,6 +770,7 @@ func (h *OpenAIGatewayHandler) Messages(c *gin.Context) {
service.OpenAIUpstreamTransportAny,
service.OpenAIEndpointCapabilityChatCompletions,
false,
requestPlatform,
)
if err != nil {
reqLog.Warn("openai_messages.account_select_failed",
@@ -1322,6 +1333,7 @@ func (h *OpenAIGatewayHandler) ResponsesWebSocket(c *gin.Context) {
}
subscription, _ := middleware2.GetSubscriptionFromContext(c)
requestPlatform := openAICompatibleRequestPlatform(apiKey)
if err := h.billingCacheService.CheckBillingEligibility(ctx, apiKey.User, apiKey, apiKey.Group, subscription, service.QuotaPlatform(c.Request.Context(), apiKey)); err != nil {
reqLog.Info("openai.websocket_billing_eligibility_check_failed", zap.Error(err))
closeOpenAIClientWS(wsConn, coderws.StatusPolicyViolation, "billing check failed")
@@ -1350,6 +1362,7 @@ func (h *OpenAIGatewayHandler) ResponsesWebSocket(c *gin.Context) {
service.OpenAIUpstreamTransportResponsesWebsocketV2,
service.OpenAIEndpointCapabilityChatCompletions,
false,
requestPlatform,
)
if err != nil {
reqLog.Warn("openai.websocket_account_select_failed",
@@ -1341,6 +1341,7 @@ func TestOpenAIResponsesWebSocket_FailoverOnUpstreamUsageLimitEvent(t *testing.T
nil,
nil,
nil,
nil,
)
cache := &concurrencyCacheMock{
@@ -1523,6 +1524,7 @@ func runOpenAIResponsesWebSocketUsageLogCase(t *testing.T, tc openAIResponsesWSU
&service.DeferredService{},
nil,
nil,
nil,
channelSvc,
nil,
nil,
@@ -136,6 +136,7 @@ func TestOpenAIGatewayHandlerImages_ServerErrorFailsOverAndReturnsClearErrorWhen
nil,
nil,
nil,
nil,
)
billingService := service.NewBillingCacheService(nil, nil, nil, nil, nil, nil, cfg, nil)
t.Cleanup(billingService.Stop)
+3
View File
@@ -20,6 +20,7 @@ func ProvideAdminHandlers(
openaiOAuthHandler *admin.OpenAIOAuthHandler,
geminiOAuthHandler *admin.GeminiOAuthHandler,
antigravityOAuthHandler *admin.AntigravityOAuthHandler,
grokOAuthHandler *admin.GrokOAuthHandler,
proxyHandler *admin.ProxyHandler,
redeemHandler *admin.RedeemHandler,
promoHandler *admin.PromoHandler,
@@ -53,6 +54,7 @@ func ProvideAdminHandlers(
OpenAIOAuth: openaiOAuthHandler,
GeminiOAuth: geminiOAuthHandler,
AntigravityOAuth: antigravityOAuthHandler,
GrokOAuth: grokOAuthHandler,
Proxy: proxyHandler,
Redeem: redeemHandler,
Promo: promoHandler,
@@ -167,6 +169,7 @@ var ProviderSet = wire.NewSet(
admin.NewOpenAIOAuthHandler,
admin.NewGeminiOAuthHandler,
admin.NewAntigravityOAuthHandler,
admin.NewGrokOAuthHandler,
admin.NewProxyHandler,
admin.NewRedeemHandler,
admin.NewPromoHandler,