Merge pull request #3558 from heathermhuang/codex/grok-cli-compat

fix: support Grok admin test and CLI compatibility
This commit is contained in:
Wesley Liddick
2026-06-30 08:42:59 +08:00
committed by GitHub
19 changed files with 901 additions and 74 deletions
+5 -2
View File
@@ -614,9 +614,12 @@ Sub2API supports Grok subscription accounts through xAI OAuth and forwards OpenA
- Platform name: `grok`
- Account type: OAuth subscription accounts
- Public gateway target: `/v1/responses` and `/responses`, forwarded to `${XAI_BASE_URL:-https://api.x.ai/v1}/responses`
- Public Responses targets: `/v1/responses`, `/responses`, and `/backend-api/codex/responses`, forwarded to `${XAI_BASE_URL:-https://api.x.ai/v1}/responses`
- Public Claude-compatible target: `/v1/messages`, converted to xAI Responses and returned as Anthropic Messages output for Claude CLI style clients
- Public Chat Completions targets: `/v1/chat/completions` and `/chat/completions`, forwarded to `${XAI_BASE_URL:-https://api.x.ai/v1}/chat/completions`
- Codex CLI style Responses WebSocket ingress is accepted on the Responses targets and bridged to xAI HTTP/SSE Responses upstream
- Initial models: `grok-4.3`, `grok-build-0.1`, `grok-4.20-0309-reasoning`, `grok-4.20-0309-non-reasoning`, and `grok-4.20-multi-agent-0309`
- Out of scope for this provider: public Grok Chat Completions routes, image, video, TTS, transcription, browser automation, cookies, and Grok web scraping
- Out of scope for this provider: image, video, TTS, transcription, browser automation, cookies, and Grok web scraping
### OAuth Configuration
+1 -1
View File
@@ -184,7 +184,7 @@ func initializeApplication(buildInfo handler.BuildInfo) (*Application, error) {
grokQuotaFetcher := service.NewGrokQuotaFetcher()
usageCache := service.NewUsageCache()
accountUsageService := service.NewAccountUsageService(accountRepository, usageLogRepository, claudeUsageFetcher, geminiQuotaService, antigravityQuotaFetcher, grokQuotaFetcher, usageCache, identityCache, tlsFingerprintProfileService)
accountTestService := service.NewAccountTestService(accountRepository, geminiTokenProvider, claudeTokenProvider, antigravityGatewayService, httpUpstream, configConfig, tlsFingerprintProfileService)
accountTestService := service.NewAccountTestService(accountRepository, geminiTokenProvider, claudeTokenProvider, grokTokenProvider, 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)
adminAnnouncementHandler := admin.NewAnnouncementHandler(announcementService)
@@ -11,6 +11,7 @@ import (
"log"
"log/slog"
"net/http"
"sort"
"strconv"
"strings"
"sync"
@@ -25,6 +26,7 @@ import (
"github.com/Wei-Shaw/sub2api/internal/pkg/openai"
"github.com/Wei-Shaw/sub2api/internal/pkg/response"
"github.com/Wei-Shaw/sub2api/internal/pkg/timezone"
"github.com/Wei-Shaw/sub2api/internal/pkg/xai"
"github.com/Wei-Shaw/sub2api/internal/service"
"github.com/gin-gonic/gin"
@@ -2064,6 +2066,56 @@ func (h *AccountHandler) GetAvailableModels(c *gin.Context) {
return
}
// Handle Grok accounts
if account.Platform == service.PlatformGrok {
defaultModels := xai.DefaultModels()
hasExplicitMapping := false
switch rawMapping := account.Credentials["model_mapping"].(type) {
case map[string]any:
hasExplicitMapping = len(rawMapping) > 0
case map[string]string:
hasExplicitMapping = len(rawMapping) > 0
}
if !hasExplicitMapping {
response.Success(c, defaultModels)
return
}
mapping := account.GetModelMapping()
if len(mapping) == 0 {
response.Success(c, defaultModels)
return
}
defaultByID := make(map[string]xai.Model, len(defaultModels))
for _, model := range defaultModels {
defaultByID[model.ID] = model
}
requestedModels := make([]string, 0, len(mapping))
for requestedModel := range mapping {
requestedModels = append(requestedModels, requestedModel)
}
sort.Strings(requestedModels)
var models []xai.Model
for _, requestedModel := range requestedModels {
if defaultModel, found := defaultByID[requestedModel]; found {
models = append(models, defaultModel)
continue
}
models = append(models, xai.Model{
ID: requestedModel,
Object: "model",
OwnedBy: "xai",
DisplayName: requestedModel,
})
}
response.Success(c, models)
return
}
// Handle Claude/Anthropic accounts
// For OAuth and Setup-Token accounts: return default models
if account.IsOAuth() {
@@ -61,6 +61,7 @@ func setupSyncUpstreamModelsRouter(adminSvc service.AdminService, upstream servi
nil,
nil,
nil,
nil,
upstream,
&config.Config{Security: config.SecurityConfig{URLAllowlist: config.URLAllowlistConfig{Enabled: false}}},
nil,
@@ -70,6 +71,77 @@ func setupSyncUpstreamModelsRouter(adminSvc service.AdminService, upstream servi
return router
}
func TestAccountHandlerGetAvailableModels_GrokUsesXAIModels(t *testing.T) {
svc := &availableModelsAdminService{
stubAdminService: newStubAdminService(),
account: service.Account{
ID: 44,
Name: "grok-oauth",
Platform: service.PlatformGrok,
Type: service.AccountTypeOAuth,
Status: service.StatusActive,
Credentials: map[string]any{
"model_mapping": map[string]any{
"grok-4.3": "grok-4.3",
},
},
},
}
router := setupAvailableModelsRouter(svc)
rec := httptest.NewRecorder()
req := httptest.NewRequest(http.MethodGet, "/api/v1/admin/accounts/44/models", nil)
router.ServeHTTP(rec, req)
require.Equal(t, http.StatusOK, rec.Code)
var resp struct {
Data []struct {
ID string `json:"id"`
} `json:"data"`
}
require.NoError(t, json.Unmarshal(rec.Body.Bytes(), &resp))
require.Len(t, resp.Data, 1)
require.Equal(t, "grok-4.3", resp.Data[0].ID)
}
func TestAccountHandlerGetAvailableModels_GrokDefaultsToXAIModelsWithoutMapping(t *testing.T) {
svc := &availableModelsAdminService{
stubAdminService: newStubAdminService(),
account: service.Account{
ID: 45,
Name: "grok-oauth-defaults",
Platform: service.PlatformGrok,
Type: service.AccountTypeOAuth,
Status: service.StatusActive,
},
}
router := setupAvailableModelsRouter(svc)
rec := httptest.NewRecorder()
req := httptest.NewRequest(http.MethodGet, "/api/v1/admin/accounts/45/models", nil)
router.ServeHTTP(rec, req)
require.Equal(t, http.StatusOK, rec.Code)
var resp struct {
Data []struct {
ID string `json:"id"`
} `json:"data"`
}
require.NoError(t, json.Unmarshal(rec.Body.Bytes(), &resp))
require.NotEmpty(t, resp.Data)
var ids []string
for _, model := range resp.Data {
id := model.ID
ids = append(ids, id)
require.NotContains(t, strings.ToLower(id), "claude")
}
require.Contains(t, ids, "grok-4.3")
require.Contains(t, ids, "grok-build-0.1")
}
func TestAccountHandlerGetAvailableModels_OpenAIOAuthUsesExplicitModelMapping(t *testing.T) {
svc := &availableModelsAdminService{
stubAdminService: newStubAdminService(),
@@ -104,6 +104,16 @@ func openAICompatibleRequestPlatform(apiKey *service.APIKey) string {
return service.PlatformOpenAI
}
func allowOpenAICompatibleMessagesDispatch(apiKey *service.APIKey) bool {
if apiKey == nil || apiKey.Group == nil {
return true
}
if apiKey.Group.Platform == service.PlatformGrok {
return true
}
return apiKey.Group.AllowMessagesDispatch
}
// NewOpenAIGatewayHandler creates a new OpenAIGatewayHandler
func NewOpenAIGatewayHandler(
gatewayService *service.OpenAIGatewayService,
@@ -660,7 +670,7 @@ func (h *OpenAIGatewayHandler) Messages(c *gin.Context) {
)
// 检查分组是否允许 /v1/messages 调度
if apiKey.Group != nil && !apiKey.Group.AllowMessagesDispatch {
if !allowOpenAICompatibleMessagesDispatch(apiKey) {
h.anthropicErrorResponse(c, http.StatusForbidden, "permission_error",
"This group does not allow /v1/messages dispatch")
return
@@ -1338,6 +1348,10 @@ func (h *OpenAIGatewayHandler) ResponsesWebSocket(c *gin.Context) {
subscription, _ := middleware2.GetSubscriptionFromContext(c)
requestPlatform := openAICompatibleRequestPlatform(apiKey)
requiredTransport := service.OpenAIUpstreamTransportResponsesWebsocketV2
if requestPlatform == service.PlatformGrok {
requiredTransport = service.OpenAIUpstreamTransportHTTPSSE
}
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")
@@ -1363,7 +1377,7 @@ func (h *OpenAIGatewayHandler) ResponsesWebSocket(c *gin.Context) {
sessionHash,
reqModel,
failedAccountIDs,
service.OpenAIUpstreamTransportResponsesWebsocketV2,
requiredTransport,
service.OpenAIEndpointCapabilityChatCompletions,
false,
requestPlatform,
@@ -434,6 +434,16 @@ func TestResolveOpenAIMessagesDispatchMappedModel(t *testing.T) {
require.Empty(t, resolveOpenAIMessagesDispatchMappedModel(&service.APIKey{Group: &service.Group{}}, "gpt-5.4"))
})
t.Run("grok_group_maps_claude_cli_model_to_grok_default", func(t *testing.T) {
apiKey := &service.APIKey{
Group: &service.Group{
Platform: service.PlatformGrok,
},
}
require.Equal(t, "grok-4.3", resolveOpenAIMessagesDispatchMappedModel(apiKey, "claude-sonnet-4-5"))
require.Empty(t, resolveOpenAIMessagesDispatchMappedModel(apiKey, "grok"))
})
t.Run("does_not_fall_back_to_group_default_mapped_model", func(t *testing.T) {
apiKey := &service.APIKey{
Group: &service.Group{
@@ -445,6 +455,60 @@ func TestResolveOpenAIMessagesDispatchMappedModel(t *testing.T) {
})
}
func TestOpenAIGatewayMessagesDispatchGateAllowsGrokGroups(t *testing.T) {
gin.SetMode(gin.TestMode)
t.Run("openai_group_without_dispatch_flag_is_rejected", func(t *testing.T) {
rec := httptest.NewRecorder()
c, _ := gin.CreateTestContext(rec)
c.Request = httptest.NewRequest(http.MethodPost, "/v1/messages", strings.NewReader(`{"model":"claude-sonnet-4-5","messages":[{"role":"user","content":"hi"}]}`))
groupID := int64(4101)
c.Set(string(middleware.ContextKeyAPIKey), &service.APIKey{
ID: 5101,
GroupID: &groupID,
User: &service.User{ID: 6101},
Group: &service.Group{
ID: groupID,
Platform: service.PlatformOpenAI,
AllowMessagesDispatch: false,
},
})
c.Set(string(middleware.ContextKeyUser), middleware.AuthSubject{UserID: 6101, Concurrency: 1})
h := &OpenAIGatewayHandler{}
h.Messages(c)
require.Equal(t, http.StatusForbidden, rec.Code)
require.Equal(t, "permission_error", gjson.GetBytes(rec.Body.Bytes(), "error.type").String())
require.Contains(t, rec.Body.String(), "This group does not allow /v1/messages dispatch")
})
t.Run("grok_group_without_dispatch_flag_reaches_gateway_dependencies", func(t *testing.T) {
rec := httptest.NewRecorder()
c, _ := gin.CreateTestContext(rec)
c.Request = httptest.NewRequest(http.MethodPost, "/v1/messages", strings.NewReader(`{"model":"grok-4.3","messages":[{"role":"user","content":"hi"}]}`))
groupID := int64(4102)
c.Set(string(middleware.ContextKeyAPIKey), &service.APIKey{
ID: 5102,
GroupID: &groupID,
User: &service.User{ID: 6102},
Group: &service.Group{
ID: groupID,
Platform: service.PlatformGrok,
AllowMessagesDispatch: false,
},
})
c.Set(string(middleware.ContextKeyUser), middleware.AuthSubject{UserID: 6102, Concurrency: 1})
h := &OpenAIGatewayHandler{}
h.Messages(c)
require.Equal(t, http.StatusServiceUnavailable, rec.Code)
require.Equal(t, "api_error", gjson.GetBytes(rec.Body.Bytes(), "error.type").String())
require.NotContains(t, rec.Body.String(), "This group does not allow /v1/messages dispatch")
})
}
func TestOpenAIModelMappedBody(t *testing.T) {
body := []byte(`{"model":"alias","input":"hello"}`)
calls := 0
+3 -37
View File
@@ -42,16 +42,6 @@ func RegisterGatewayRoutes(
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)
@@ -63,11 +53,7 @@ func RegisterGatewayRoutes(
{
// /v1/messages: auto-route based on group platform
gateway.POST("/messages", func(c *gin.Context) {
if getGroupPlatform(c) == service.PlatformGrok {
rejectGrokUnsupportedEndpoint(c, "Messages API")
return
}
if isOpenAIGatewayPlatform(c) {
if isOpenAIResponsesCompatibleGatewayPlatform(c) {
h.OpenAIGateway.Messages(c)
return
}
@@ -111,19 +97,11 @@ func RegisterGatewayRoutes(
h.Gateway.Responses(c)
})
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.PlatformGrok {
rejectGrokUnsupportedEndpoint(c, "Chat Completions API")
return
}
if isOpenAIGatewayPlatform(c) {
if isOpenAIResponsesCompatibleGatewayPlatform(c) {
h.OpenAIGateway.ChatCompletions(c)
return
}
@@ -196,10 +174,6 @@ 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, 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")
@@ -208,20 +182,12 @@ func RegisterGatewayRoutes(
codexDirect.POST("/responses", responsesHandler)
codexDirect.POST("/responses/*subpath", responsesHandler)
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.PlatformGrok {
rejectGrokUnsupportedEndpoint(c, "Chat Completions API")
return
}
if isOpenAIGatewayPlatform(c) {
if isOpenAIResponsesCompatibleGatewayPlatform(c) {
h.OpenAIGateway.ChatCompletions(c)
return
}
@@ -83,7 +83,7 @@ func TestGatewayRoutesOpenAIImagesPathsAreRegistered(t *testing.T) {
}
}
func TestGatewayRoutesGrokOnlyAllowsResponsesHTTP(t *testing.T) {
func TestGatewayRoutesGrokAllowsCLICompatibilityEntrypoints(t *testing.T) {
router := newGatewayRoutesTestRouter(service.PlatformGrok)
for _, tc := range []struct {
@@ -102,8 +102,8 @@ func TestGatewayRoutesGrokOnlyAllowsResponsesHTTP(t *testing.T) {
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")
require.NotEqual(t, http.StatusNotFound, w.Code, "method=%s path=%s", tc.method, tc.path)
require.NotContains(t, w.Body.String(), "not supported for Grok groups")
}
req := httptest.NewRequest(http.MethodPost, "/v1/messages/count_tokens", strings.NewReader(`{"model":"grok","messages":[{"role":"user","content":"hi"}]}`))
@@ -22,6 +22,7 @@ import (
"github.com/Wei-Shaw/sub2api/internal/pkg/geminicli"
"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/urlvalidator"
"github.com/gin-gonic/gin"
"github.com/google/uuid"
@@ -66,6 +67,7 @@ type AccountTestService struct {
accountRepo AccountRepository
geminiTokenProvider *GeminiTokenProvider
claudeTokenProvider *ClaudeTokenProvider
grokTokenProvider *GrokTokenProvider
antigravityGatewayService *AntigravityGatewayService
httpUpstream HTTPUpstream
cfg *config.Config
@@ -77,6 +79,7 @@ func NewAccountTestService(
accountRepo AccountRepository,
geminiTokenProvider *GeminiTokenProvider,
claudeTokenProvider *ClaudeTokenProvider,
grokTokenProvider *GrokTokenProvider,
antigravityGatewayService *AntigravityGatewayService,
httpUpstream HTTPUpstream,
cfg *config.Config,
@@ -86,6 +89,7 @@ func NewAccountTestService(
accountRepo: accountRepo,
geminiTokenProvider: geminiTokenProvider,
claudeTokenProvider: claudeTokenProvider,
grokTokenProvider: grokTokenProvider,
antigravityGatewayService: antigravityGatewayService,
httpUpstream: httpUpstream,
cfg: cfg,
@@ -188,6 +192,10 @@ func (s *AccountTestService) TestAccountConnection(c *gin.Context, accountID int
return s.testGeminiAccountConnection(c, account, modelID, prompt)
}
if account.Platform == PlatformGrok {
return s.testGrokAccountConnection(c, account, modelID)
}
if account.Platform == PlatformAntigravity {
return s.routeAntigravityTest(c, account, modelID, prompt)
}
@@ -627,6 +635,89 @@ func (s *AccountTestService) testOpenAIAccountConnection(c *gin.Context, account
return s.processOpenAIStream(c, resp.Body)
}
// testGrokAccountConnection tests a Grok OAuth account through xAI's Responses API.
func (s *AccountTestService) testGrokAccountConnection(c *gin.Context, account *Account, modelID string) error {
ctx := c.Request.Context()
if account.Type != AccountTypeOAuth {
return s.sendErrorAndEnd(c, fmt.Sprintf("Unsupported Grok account type: %s", account.Type))
}
if s.grokTokenProvider == nil {
return s.sendErrorAndEnd(c, "Grok token provider not configured")
}
if s.httpUpstream == nil {
return s.sendErrorAndEnd(c, "HTTP upstream not configured")
}
testModelID := strings.TrimSpace(modelID)
if testModelID == "" {
testModelID = "grok-4.3"
}
if mapped := strings.TrimSpace(account.GetMappedModel(testModelID)); mapped != "" {
testModelID = mapped
}
authToken, err := s.grokTokenProvider.GetAccessToken(ctx, account)
if err != nil {
return s.sendErrorAndEnd(c, fmt.Sprintf("Failed to get Grok access token: %s", err.Error()))
}
apiURL, err := xai.BuildResponsesURL(account.GetGrokBaseURL())
if err != nil {
return s.sendErrorAndEnd(c, fmt.Sprintf("Invalid Grok base URL: %s", err.Error()))
}
c.Writer.Header().Set("Content-Type", "text/event-stream")
c.Writer.Header().Set("Cache-Control", "no-cache")
c.Writer.Header().Set("Connection", "keep-alive")
c.Writer.Header().Set("X-Accel-Buffering", "no")
c.Writer.Flush()
payloadBytes, err := json.Marshal(map[string]any{
"model": testModelID,
"input": "hi",
"stream": true,
})
if err != nil {
return s.sendErrorAndEnd(c, "Failed to create Grok test payload")
}
s.sendEvent(c, TestEvent{Type: "test_start", Model: testModelID})
req, err := http.NewRequestWithContext(ctx, http.MethodPost, apiURL, bytes.NewReader(payloadBytes))
if err != nil {
return s.sendErrorAndEnd(c, "Failed to create Grok request")
}
req.Header.Set("Content-Type", "application/json")
req.Header.Set("Accept", "application/json, text/event-stream")
req.Header.Set("Authorization", "Bearer "+authToken)
req.Header.Set("User-Agent", "sub2api-grok/1.0")
proxyURL := ""
if account.ProxyID != nil && account.Proxy != nil {
proxyURL = account.Proxy.URL()
}
resp, err := s.httpUpstream.Do(req, proxyURL, account.ID, account.Concurrency)
if err != nil {
return s.sendErrorAndEnd(c, fmt.Sprintf("Grok Responses API request failed: %s", err.Error()))
}
defer func() { _ = resp.Body.Close() }()
if snapshot := xai.ParseQuotaHeaders(resp.Header, resp.StatusCode); snapshot != nil && s.accountRepo != nil {
_ = s.accountRepo.UpdateExtra(ctx, account.ID, map[string]any{
grokQuotaSnapshotExtraKey: snapshot,
})
}
if resp.StatusCode != http.StatusOK {
body, _ := io.ReadAll(resp.Body)
return s.sendErrorAndEnd(c, fmt.Sprintf("Grok Responses API returned %d: %s", resp.StatusCode, string(body)))
}
return s.processOpenAIStream(c, resp.Body)
}
// testOpenAIChatCompletionsConnection tests an OpenAI-compatible APIKey account
// through the raw /v1/chat/completions endpoint.
func (s *AccountTestService) testOpenAIChatCompletionsConnection(
@@ -0,0 +1,67 @@
//go:build unit
package service
import (
"io"
"net/http"
"net/http/httptest"
"strings"
"testing"
"time"
"github.com/gin-gonic/gin"
"github.com/stretchr/testify/require"
"github.com/tidwall/gjson"
)
func TestAccountTestService_TestAccountConnection_GrokUsesXAIResponses(t *testing.T) {
gin.SetMode(gin.TestMode)
account := &Account{
ID: 13,
Name: "grok-oauth",
Platform: PlatformGrok,
Type: AccountTypeOAuth,
Status: StatusActive,
Schedulable: true,
Concurrency: 1,
Credentials: map[string]any{
"access_token": "grok-access-token",
"expires_at": time.Now().Add(time.Hour).UTC().Format(time.RFC3339),
"model_mapping": map[string]any{
"grok": "grok-4.3",
},
},
}
repo := &mockAccountRepoForGemini{
accountsByID: map[int64]*Account{account.ID: account},
}
upstream := &httpUpstreamRecorder{resp: &http.Response{
StatusCode: http.StatusOK,
Header: http.Header{"Content-Type": []string{"text/event-stream"}},
Body: io.NopCloser(strings.NewReader(
"data: {\"type\":\"response.output_text.delta\",\"delta\":\"ok\"}\n\n" +
"data: {\"type\":\"response.completed\"}\n\n",
)),
}}
svc := &AccountTestService{
accountRepo: repo,
grokTokenProvider: NewGrokTokenProvider(repo, nil),
httpUpstream: upstream,
}
rec := httptest.NewRecorder()
c, _ := gin.CreateTestContext(rec)
c.Request = httptest.NewRequest(http.MethodPost, "/api/v1/admin/accounts/13/test", nil)
err := svc.TestAccountConnection(c, account.ID, "grok", "", AccountTestModeDefault)
require.NoError(t, err)
require.Equal(t, "https://api.x.ai/v1/responses", upstream.lastReq.URL.String())
require.Equal(t, "Bearer grok-access-token", upstream.lastReq.Header.Get("Authorization"))
require.Equal(t, "grok-4.3", gjson.GetBytes(upstream.lastBody, "model").String())
require.NotContains(t, rec.Body.String(), "claude")
require.Contains(t, rec.Body.String(), `"model":"grok-4.3"`)
require.Contains(t, rec.Body.String(), `"type":"test_complete"`)
}
@@ -147,9 +147,168 @@ func patchGrokResponsesBody(body []byte, upstreamModel string) ([]byte, error) {
}
}
}
out, err = sanitizeGrokResponsesUnsupportedFields(out)
if err != nil {
return nil, err
}
out, err = sanitizeGrokResponsesTools(out)
if err != nil {
return nil, err
}
return out, nil
}
var grokResponsesUnsupportedRecursiveFields = map[string]struct{}{
"external_web_access": {},
}
func sanitizeGrokResponsesUnsupportedFields(body []byte) ([]byte, error) {
if !bytes.Contains(body, []byte(`"external_web_access"`)) {
return body, nil
}
var payload any
if err := json.Unmarshal(body, &payload); err != nil {
return nil, err
}
if !deleteJSONFields(payload, grokResponsesUnsupportedRecursiveFields) {
return body, nil
}
return json.Marshal(payload)
}
func deleteJSONFields(value any, fields map[string]struct{}) bool {
switch typed := value.(type) {
case map[string]any:
changed := false
for field := range fields {
if _, ok := typed[field]; ok {
delete(typed, field)
changed = true
}
}
for _, child := range typed {
if deleteJSONFields(child, fields) {
changed = true
}
}
return changed
case []any:
changed := false
for _, child := range typed {
if deleteJSONFields(child, fields) {
changed = true
}
}
return changed
default:
return false
}
}
var grokResponsesSupportedToolTypes = map[string]struct{}{
"code_execution": {},
"code_interpreter": {},
"collections_search": {},
"file_search": {},
"function": {},
"mcp": {},
"shell": {},
"web_search": {},
"x_search": {},
}
func sanitizeGrokResponsesTools(body []byte) ([]byte, error) {
tools := gjson.GetBytes(body, "tools")
if !tools.Exists() || !tools.IsArray() {
return body, nil
}
rawTools := tools.Array()
filteredTools := make([]json.RawMessage, 0, len(rawTools))
for _, tool := range rawTools {
toolType := strings.TrimSpace(tool.Get("type").String())
if _, ok := grokResponsesSupportedToolTypes[toolType]; ok {
filteredTools = append(filteredTools, json.RawMessage(tool.Raw))
}
}
var err error
if len(filteredTools) != len(rawTools) {
if len(filteredTools) == 0 {
body, err = sjson.DeleteBytes(body, "tools")
} else {
var encoded []byte
encoded, err = json.Marshal(filteredTools)
if err != nil {
return nil, err
}
body, err = sjson.SetRawBytes(body, "tools", encoded)
}
if err != nil {
return nil, err
}
}
toolChoice := gjson.GetBytes(body, "tool_choice")
if !toolChoice.Exists() {
return body, nil
}
if shouldDropGrokToolChoice(toolChoice, filteredTools) {
body, err = sjson.DeleteBytes(body, "tool_choice")
if err != nil {
return nil, err
}
}
return body, nil
}
func shouldDropGrokToolChoice(toolChoice gjson.Result, tools []json.RawMessage) bool {
if len(tools) == 0 {
return true
}
if !toolChoice.IsObject() {
return false
}
choiceType := strings.TrimSpace(toolChoice.Get("type").String())
if choiceType == "" {
return false
}
if _, ok := grokResponsesSupportedToolTypes[choiceType]; !ok {
return true
}
if choiceType == "function" {
choiceName := strings.TrimSpace(toolChoice.Get("name").String())
if choiceName == "" {
choiceName = strings.TrimSpace(toolChoice.Get("function.name").String())
}
if choiceName == "" {
return false
}
for _, tool := range tools {
var item struct {
Type string `json:"type"`
Name string `json:"name"`
Function struct {
Name string `json:"name"`
} `json:"function"`
}
if err := json.Unmarshal(tool, &item); err != nil {
continue
}
name := strings.TrimSpace(item.Name)
if name == "" {
name = strings.TrimSpace(item.Function.Name)
}
if strings.TrimSpace(item.Type) == "function" && name == choiceName {
return false
}
}
return true
}
return false
}
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 {
@@ -39,6 +39,71 @@ func TestPatchGrokResponsesBodySetsMappedModelAndDropsUnsupportedFields(t *testi
require.Equal(t, "high", gjson.GetBytes(patched, "reasoning.effort").String())
}
func TestPatchGrokResponsesBodyDropsNestedUnsupportedFields(t *testing.T) {
t.Parallel()
body := []byte(`{
"model": "grok",
"input": "hello",
"external_web_access": true,
"tools": [
{"type": "function", "name": "kept_fn", "external_web_access": true, "parameters": {"type": "object", "properties": {"q": {"type": "string", "external_web_access": true}}}}
],
"metadata": {"external_web_access": false}
}`)
patched, err := patchGrokResponsesBody(body, "grok-4.3")
require.NoError(t, err)
require.True(t, json.Valid(patched))
require.False(t, strings.Contains(string(patched), "external_web_access"))
require.Equal(t, "kept_fn", gjson.GetBytes(patched, "tools.0.name").String())
}
func TestPatchGrokResponsesBodyDropsUnsupportedNamespaceTools(t *testing.T) {
t.Parallel()
body := []byte(`{
"model": "grok",
"input": "hello",
"tools": [
{"type": "namespace", "namespace": "functions", "tools": [{"type": "function", "name": "inner"}]},
{"type": "function", "name": "kept_fn", "parameters": {"type": "object"}},
{"type": "shell", "name": "kept_shell"}
],
"tool_choice": {"type": "function", "name": "kept_fn"}
}`)
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.Len(t, gjson.GetBytes(patched, "tools").Array(), 2)
require.False(t, gjson.GetBytes(patched, `tools.#(type=="namespace")`).Exists())
require.True(t, gjson.GetBytes(patched, `tools.#(type=="function")`).Exists())
require.True(t, gjson.GetBytes(patched, `tools.#(type=="shell")`).Exists())
require.Equal(t, "kept_fn", gjson.GetBytes(patched, "tool_choice.name").String())
}
func TestPatchGrokResponsesBodyDropsToolChoiceWhenNoSupportedToolsRemain(t *testing.T) {
t.Parallel()
body := []byte(`{
"model": "grok",
"input": "hello",
"tools": [
{"type": "namespace", "namespace": "functions"},
{"type": "image_generation", "model": "gpt-image-2"}
],
"tool_choice": {"type": "namespace", "namespace": "functions"}
}`)
patched, err := patchGrokResponsesBody(body, "grok-4.3")
require.NoError(t, err)
require.True(t, json.Valid(patched))
require.False(t, gjson.GetBytes(patched, "tools").Exists())
require.False(t, gjson.GetBytes(patched, "tool_choice").Exists())
}
func TestBuildGrokResponsesRequestUsesAccountBaseURLAndBearerToken(t *testing.T) {
t.Setenv(xai.EnvAllowUnsafeURLOverrides, "true")
@@ -271,6 +336,54 @@ func TestForwardAsChatCompletionsForGrokStreamingUsesRawXAIChatCompletions(t *te
require.NotNil(t, repo.updates[53][grokQuotaSnapshotExtraKey])
}
func TestForwardAsAnthropicForGrokUsesXAIResponses(t *testing.T) {
gin.SetMode(gin.TestMode)
recorder := httptest.NewRecorder()
c, _ := gin.CreateTestContext(recorder)
body := []byte(`{"model":"grok","max_tokens":32,"stream":false,"messages":[{"role":"user","content":"hi"}]}`)
c.Request = httptest.NewRequest(http.MethodPost, "/v1/messages", bytes.NewReader(body))
account := &Account{
ID: 54,
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{54: account},
},
}
upstream := &httpUpstreamRecorder{resp: openAICompatSSECompletedResponse("resp_grok_messages", "grok-4.3")}
svc := &OpenAIGatewayService{
httpUpstream: upstream,
grokTokenProvider: NewGrokTokenProvider(repo, nil),
accountRepo: repo,
}
result, err := svc.ForwardAsAnthropic(context.Background(), c, account, body, "", "")
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, "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").Bool())
require.NotContains(t, string(upstream.lastBody), "chatgpt.com")
require.Equal(t, "grok", result.Model)
require.Equal(t, "grok-4.3", result.UpstreamModel)
require.Equal(t, 5, result.Usage.InputTokens)
require.Equal(t, 2, result.Usage.OutputTokens)
require.Contains(t, recorder.Body.String(), `"type":"message"`)
require.Contains(t, recorder.Body.String(), "ok")
}
func TestHandleGrokAccountUpstreamErrorTempUnschedulesReadinessStates(t *testing.T) {
tests := []struct {
name string
@@ -16,6 +16,7 @@ import (
"github.com/Wei-Shaw/sub2api/internal/pkg/apicompat"
"github.com/Wei-Shaw/sub2api/internal/pkg/claude"
"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"
"go.uber.org/zap"
@@ -149,7 +150,7 @@ func (s *OpenAIGatewayService) ForwardAsAnthropic(
return nil, fmt.Errorf("marshal responses request: %w", err)
}
if account.Type == AccountTypeOAuth {
if account.Type == AccountTypeOAuth && account.Platform != PlatformGrok {
var reqBody map[string]any
if err := json.Unmarshal(responsesBody, &reqBody); err != nil {
return nil, fmt.Errorf("unmarshal for codex transform: %w", err)
@@ -237,6 +238,13 @@ func (s *OpenAIGatewayService) ForwardAsAnthropic(
return nil, policyErr
}
responsesBody = updatedBody
if account.Platform == PlatformGrok {
patchedBody, patchErr := patchGrokResponsesBody(responsesBody, upstreamModel)
if patchErr != nil {
return nil, patchErr
}
responsesBody = patchedBody
}
// 5. Get access token
token, _, err := s.GetAccessToken(ctx, account)
@@ -246,7 +254,12 @@ func (s *OpenAIGatewayService) ForwardAsAnthropic(
// 6. Build upstream request
upstreamCtx, releaseUpstreamCtx := detachUpstreamContext(ctx)
upstreamReq, err := s.buildUpstreamRequest(upstreamCtx, c, account, responsesBody, token, isStream, promptCacheKey, false)
var upstreamReq *http.Request
if account.Platform == PlatformGrok {
upstreamReq, err = buildGrokResponsesRequest(upstreamCtx, c, account, responsesBody, token)
} else {
upstreamReq, err = s.buildUpstreamRequest(upstreamCtx, c, account, responsesBody, token, isStream, promptCacheKey, false)
}
releaseUpstreamCtx()
if err != nil {
return nil, fmt.Errorf("build upstream request: %w", err)
@@ -261,7 +274,7 @@ func (s *OpenAIGatewayService) ForwardAsAnthropic(
upstreamReq.Header.Set("conversation_id", isolatedSessionID)
}
}
if account.Type == AccountTypeOAuth {
if account.Type == AccountTypeOAuth && account.Platform != PlatformGrok {
// Anthropic Messages compatibility uses the ChatGPT Codex SSE endpoint.
// Match airgate-openai's request shape: the SSE endpoint does not need
// the Responses experimental beta header, and forcing originator can make
@@ -303,6 +316,10 @@ func (s *OpenAIGatewayService) ForwardAsAnthropic(
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))
s.handleGrokAccountUpstreamError(ctx, account, resp.StatusCode, resp.Header, respBody)
}
upstreamMsg := strings.TrimSpace(extractUpstreamErrorMessage(respBody))
upstreamMsg = sanitizeUpstreamErrorMessage(upstreamMsg)
@@ -395,7 +412,9 @@ func (s *OpenAIGatewayService) ForwardAsAnthropic(
// Extract and save Codex usage snapshot from response headers (for OAuth accounts)
if handleErr == nil && account.Type == AccountTypeOAuth {
if snapshot := ParseCodexRateLimitHeaders(resp.Header); snapshot != nil {
if account.Platform == PlatformGrok {
s.updateGrokUsageSnapshot(ctx, account.ID, xai.ParseQuotaHeaders(resp.Header, resp.StatusCode))
} else if snapshot := ParseCodexRateLimitHeaders(resp.Header); snapshot != nil {
s.updateCodexUsageSnapshot(ctx, account.ID, snapshot)
}
}
@@ -1,6 +1,10 @@
package service
import "strings"
import (
"strings"
"github.com/Wei-Shaw/sub2api/internal/pkg/xai"
)
const (
defaultOpenAIMessagesDispatchOpusMappedModel = "gpt-5.4"
@@ -64,6 +68,13 @@ func (g *Group) ResolveMessagesDispatchModel(requestedModel string) string {
return ""
}
if g.Platform == PlatformGrok {
if claudeMessagesDispatchFamily(requestedModel) != "" {
return xai.DefaultModelMapping()["grok"]
}
return ""
}
cfg := normalizeOpenAIMessagesDispatchModelConfig(g.MessagesDispatchModelConfig)
if mappedModel := strings.TrimSpace(cfg.ExactModelMappings[requestedModel]); mappedModel != "" {
return mappedModel
@@ -25,3 +25,15 @@ func TestNormalizeOpenAIMessagesDispatchModelConfig(t *testing.T) {
"claude-sonnet-4-5-20250929": "gpt-5.2",
}, cfg.ExactModelMappings)
}
func TestGroupResolveMessagesDispatchModel_GrokMapsClaudeFamilyToGrok(t *testing.T) {
t.Parallel()
group := &Group{Platform: PlatformGrok}
require.Equal(t, "grok-4.3", group.ResolveMessagesDispatchModel("claude-sonnet-4-5"))
require.Equal(t, "grok-4.3", group.ResolveMessagesDispatchModel("claude-opus-4-6"))
require.Equal(t, "grok-4.3", group.ResolveMessagesDispatchModel("claude-haiku-4-5"))
require.Empty(t, group.ResolveMessagesDispatchModel("grok"))
require.Empty(t, group.ResolveMessagesDispatchModel("gpt-5.3-codex"))
}
+18 -10
View File
@@ -2491,9 +2491,10 @@ func (s *OpenAIGatewayService) ProxyResponsesWebSocketFromClient(
}
wsDecision := s.getOpenAIWSProtocolResolver().Resolve(account)
forceHTTPBridge := account.Platform == PlatformGrok
modeRouterV2Enabled := s != nil && s.cfg != nil && s.cfg.Gateway.OpenAIWS.ModeRouterV2Enabled
ingressMode := OpenAIWSIngressModeCtxPool
if modeRouterV2Enabled {
if modeRouterV2Enabled && !forceHTTPBridge {
ingressMode = account.ResolveOpenAIResponsesWebSocketV2Mode(s.cfg.Gateway.OpenAIWS.IngressModeDefault)
if ingressMode == OpenAIWSIngressModeOff {
return NewOpenAIWSClientCloseError(
@@ -2527,20 +2528,27 @@ func (s *OpenAIGatewayService) ProxyResponsesWebSocketFromClient(
)
}
}
if wsDecision.Transport != OpenAIUpstreamTransportResponsesWebsocketV2 {
if !forceHTTPBridge && wsDecision.Transport != OpenAIUpstreamTransportResponsesWebsocketV2 {
return fmt.Errorf("websocket ingress requires ws_v2 transport, got=%s", wsDecision.Transport)
}
dedicatedMode := modeRouterV2Enabled && ingressMode == OpenAIWSIngressModeDedicated
wsURL, err := s.buildOpenAIResponsesWSURL(account)
if err != nil {
return fmt.Errorf("build ws url: %w", err)
}
wsURL := ""
wsHost := "-"
wsPath := "-"
if parsedURL, parseErr := url.Parse(wsURL); parseErr == nil && parsedURL != nil {
wsHost = normalizeOpenAIWSLogValue(parsedURL.Host)
wsPath = normalizeOpenAIWSLogValue(parsedURL.Path)
if forceHTTPBridge {
wsHost = "xai-http-bridge"
wsPath = "/v1/responses"
} else {
var err error
wsURL, err = s.buildOpenAIResponsesWSURL(account)
if err != nil {
return fmt.Errorf("build ws url: %w", err)
}
if parsedURL, parseErr := url.Parse(wsURL); parseErr == nil && parsedURL != nil {
wsHost = normalizeOpenAIWSLogValue(parsedURL.Host)
wsPath = normalizeOpenAIWSLogValue(parsedURL.Path)
}
}
debugEnabled := isOpenAIWSModeDebugEnabled()
isCodexCLI := openai.IsCodexOfficialClientByHeaders(c.GetHeader("User-Agent"), c.GetHeader("originator")) || (s.cfg != nil && s.cfg.Gateway.ForceCodexCLI)
@@ -2826,7 +2834,7 @@ func (s *OpenAIGatewayService) ProxyResponsesWebSocketFromClient(
}
refreshIngressRouteState(firstPayload)
if s.shouldBridgeOpenAIWSHTTP(firstPayload.payloadBytes, firstPayload.previousResponseID) {
if s.shouldBridgeOpenAIWSHTTP(account, firstPayload.payloadBytes, firstPayload.previousResponseID) {
logOpenAIWSModeInfo(
"ingress_ws_http_bridge_start account_id=%d account_type=%s payload_bytes=%d threshold_bytes=%d has_session_hash=%v store_disabled=%v",
account.ID,
@@ -40,7 +40,10 @@ func (s *OpenAIGatewayService) openAIWSHTTPBridgeThresholdBytes() int64 {
return s.cfg.Gateway.OpenAIWS.HTTPBridgeThresholdBytes
}
func (s *OpenAIGatewayService) shouldBridgeOpenAIWSHTTP(payloadBytes int, previousResponseID string) bool {
func (s *OpenAIGatewayService) shouldBridgeOpenAIWSHTTP(account *Account, payloadBytes int, previousResponseID string) bool {
if account != nil && account.Platform == PlatformGrok {
return true
}
if !s.openAIWSHTTPBridgeEnabled() {
return false
}
@@ -174,7 +177,26 @@ func (s *OpenAIGatewayService) proxyOpenAIWSHTTPBridgeTurn(
}
upstreamCtx, releaseUpstreamCtx := detachUpstreamContext(ctx)
upstreamReq, err := s.buildUpstreamRequestOpenAIPassthrough(upstreamCtx, c, account, body, token)
var upstreamReq *http.Request
if account.Platform == PlatformGrok {
upstreamModel := strings.TrimSpace(gjson.GetBytes(body, "model").String())
if originalModel != "" {
if mappedModel := normalizeOpenAIModelForUpstream(account, account.GetMappedModel(originalModel)); mappedModel != "" {
upstreamModel = mappedModel
}
}
if upstreamModel == "" {
upstreamModel = "grok-4.3"
}
body, err = patchGrokResponsesBody(body, upstreamModel)
if err != nil {
releaseUpstreamCtx()
return nil, err
}
upstreamReq, err = buildGrokResponsesRequest(upstreamCtx, c, account, body, token)
} else {
upstreamReq, err = s.buildUpstreamRequestOpenAIPassthrough(upstreamCtx, c, account, body, token)
}
releaseUpstreamCtx()
if err != nil {
return nil, err
@@ -2,6 +2,7 @@ package service
import (
"context"
"errors"
"io"
"net/http"
"net/http/httptest"
@@ -10,6 +11,7 @@ import (
"time"
"github.com/Wei-Shaw/sub2api/internal/config"
"github.com/Wei-Shaw/sub2api/internal/pkg/xai"
coderws "github.com/coder/websocket"
"github.com/gin-gonic/gin"
"github.com/stretchr/testify/require"
@@ -39,12 +41,13 @@ func TestOpenAIWSHTTPBridgeDecisionKeepsSmallFramesOnWS(t *testing.T) {
},
}
require.False(t, svc.shouldBridgeOpenAIWSHTTP(99, ""))
require.True(t, svc.shouldBridgeOpenAIWSHTTP(100, ""))
require.False(t, svc.shouldBridgeOpenAIWSHTTP(1000, "resp_existing"))
require.False(t, svc.shouldBridgeOpenAIWSHTTP(nil, 99, ""))
require.True(t, svc.shouldBridgeOpenAIWSHTTP(nil, 100, ""))
require.False(t, svc.shouldBridgeOpenAIWSHTTP(nil, 1000, "resp_existing"))
svc.cfg.Gateway.OpenAIWS.HTTPBridgeEnabled = false
require.False(t, svc.shouldBridgeOpenAIWSHTTP(1000, ""))
require.False(t, svc.shouldBridgeOpenAIWSHTTP(nil, 1000, ""))
require.True(t, svc.shouldBridgeOpenAIWSHTTP(&Account{Platform: PlatformGrok}, 1, "resp_existing"))
}
func TestOpenAIWSHTTPBridgeRelaysSSEFramesAsWebSocketMessages(t *testing.T) {
@@ -173,6 +176,119 @@ func TestOpenAIWSHTTPBridgeRelaysSSEFramesAsWebSocketMessages(t *testing.T) {
require.True(t, gjson.GetBytes(upstream.lastBody, "stream").Bool())
}
func TestProxyResponsesWebSocketFromClientForGrokUsesXAIHTTPBridge(t *testing.T) {
gin.SetMode(gin.TestMode)
sseBody := strings.Join([]string{
`data: {"type":"response.created","response":{"id":"resp_grok_ws","model":"grok-4.3"}}`,
"",
`data: {"type":"response.output_text.delta","response":{"id":"resp_grok_ws"},"delta":"ok"}`,
"",
`data: {"type":"response.completed","response":{"id":"resp_grok_ws","model":"grok-4.3","usage":{"input_tokens":4,"output_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-ws-req"},
},
Body: io.NopCloser(strings.NewReader(sseBody)),
}}
svc := &OpenAIGatewayService{
cfg: &config.Config{
Gateway: config.GatewayConfig{
MaxLineSize: defaultMaxLineSize,
},
},
httpUpstream: upstream,
}
account := &Account{
ID: 71,
Name: "grok",
Platform: PlatformGrok,
Type: AccountTypeOAuth,
Concurrency: 1,
Status: StatusActive,
Credentials: map[string]any{
"base_url": xai.DefaultCLIBaseURL,
},
}
errCh := make(chan error, 1)
wsServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
conn, err := coderws.Accept(w, r, &coderws.AcceptOptions{CompressionMode: coderws.CompressionContextTakeover})
if err != nil {
errCh <- err
return
}
defer func() { _ = conn.CloseNow() }()
readCtx, cancelRead := context.WithTimeout(r.Context(), 3*time.Second)
msgType, firstMessage, err := conn.Read(readCtx)
cancelRead()
if err != nil {
errCh <- err
return
}
if msgType != coderws.MessageText {
errCh <- errors.New("first message was not text")
return
}
rec := httptest.NewRecorder()
ginCtx, _ := gin.CreateTestContext(rec)
req := r.Clone(r.Context())
req.Header = req.Header.Clone()
ginCtx.Request = req
errCh <- svc.ProxyResponsesWebSocketFromClient(r.Context(), ginCtx, conn, account, "access-token", firstMessage, nil)
}))
defer wsServer.Close()
dialCtx, cancelDial := context.WithTimeout(context.Background(), 3*time.Second)
clientConn, _, err := coderws.Dial(dialCtx, "ws"+strings.TrimPrefix(wsServer.URL, "http"), nil)
cancelDial()
require.NoError(t, err)
writeCtx, cancelWrite := context.WithTimeout(context.Background(), 3*time.Second)
err = clientConn.Write(writeCtx, coderws.MessageText, []byte(`{"type":"response.create","generate":true,"model":"grok","stream":true,"input":"hi","prompt_cache_retention":"24h"}`))
cancelWrite()
require.NoError(t, err)
readEvent := func() []byte {
readCtx, cancelRead := context.WithTimeout(context.Background(), 3*time.Second)
msgType, event, readErr := clientConn.Read(readCtx)
cancelRead()
require.NoError(t, readErr)
require.Equal(t, coderws.MessageText, msgType)
return event
}
created := readEvent()
delta := readEvent()
completed := readEvent()
require.Equal(t, "response.created", gjson.GetBytes(created, "type").String())
require.Equal(t, "response.output_text.delta", gjson.GetBytes(delta, "type").String())
require.Equal(t, "response.completed", gjson.GetBytes(completed, "type").String())
_ = clientConn.Close(coderws.StatusNormalClosure, "done")
select {
case proxyErr := <-errCh:
require.NoError(t, proxyErr)
case <-time.After(3 * time.Second):
require.Fail(t, "proxy did not finish after client close")
}
require.Equal(t, xai.DefaultCLIBaseURL+"/responses", upstream.lastReq.URL.String())
require.Equal(t, "Bearer access-token", upstream.lastReq.Header.Get("Authorization"))
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.False(t, gjson.GetBytes(upstream.lastBody, "type").Exists())
require.False(t, gjson.GetBytes(upstream.lastBody, "generate").Exists())
require.False(t, gjson.GetBytes(upstream.lastBody, "prompt_cache_retention").Exists())
}
func TestOpenAIWSHTTPBridgeAcceptsFirstFrameAboveLegacy16MiB(t *testing.T) {
gin.SetMode(gin.TestMode)
@@ -59,17 +59,17 @@ function createStreamResponse(lines: string[]) {
} as Response
}
function mountModal() {
function mountModal(account: Record<string, unknown> = {
id: 42,
name: 'Gemini Image Test',
platform: 'gemini',
type: 'apikey',
status: 'active'
}) {
return mount(AccountTestModal, {
props: {
show: false,
account: {
id: 42,
name: 'Gemini Image Test',
platform: 'gemini',
type: 'apikey',
status: 'active'
}
account
} as any,
global: {
stubs: {
@@ -144,4 +144,42 @@ describe('AccountTestModal', () => {
expect(preview.exists()).toBe(true)
expect(preview.attributes('src')).toBe('data:image/png;base64,QUJD')
})
it('grok 账号测试默认选择 Grok 模型', async () => {
getAvailableModels.mockResolvedValue([
{ id: 'grok-4.3', display_name: 'Grok 4.3' },
{ id: 'grok-build-0.1', display_name: 'Grok Build 0.1' }
])
global.fetch = vi.fn().mockResolvedValue(
createStreamResponse([
'data: {"type":"test_start","model":"grok-4.3"}\n',
'data: {"type":"content","text":"ok"}\n',
'data: {"type":"test_complete","success":true}\n'
])
) as any
const wrapper = mountModal({
id: 13,
name: 'Grok Account',
platform: 'grok',
type: 'oauth',
status: 'active'
})
await wrapper.setProps({ show: true })
await flushPromises()
const buttons = wrapper.findAll('button')
const startButton = buttons.find((button) => button.text().includes('admin.accounts.startTest'))
expect(startButton).toBeTruthy()
await startButton!.trigger('click')
await flushPromises()
expect(global.fetch).toHaveBeenCalledTimes(1)
const [, request] = (global.fetch as any).mock.calls[0]
expect(JSON.parse(request.body)).toEqual({
model_id: 'grok-4.3',
prompt: ''
})
})
})