From 4a7148e203eb7c56065cdd00e4d3f50071b36943 Mon Sep 17 00:00:00 2001 From: Heatherm Huang Date: Mon, 29 Jun 2026 17:53:19 +0800 Subject: [PATCH] fix: support grok cli compatibility routes --- README.md | 7 +- backend/cmd/server/wire_gen.go | 2 +- .../internal/handler/admin/account_handler.go | 52 ++++++++ .../account_handler_available_models_test.go | 72 ++++++++++ .../handler/openai_gateway_handler.go | 6 +- backend/internal/server/routes/gateway.go | 40 +----- .../internal/server/routes/gateway_test.go | 6 +- .../internal/service/account_test_service.go | 91 +++++++++++++ .../service/account_test_service_grok_test.go | 67 ++++++++++ .../service/openai_gateway_grok_test.go | 48 +++++++ .../service/openai_gateway_messages.go | 27 +++- .../internal/service/openai_ws_forwarder.go | 28 ++-- .../internal/service/openai_ws_http_bridge.go | 26 +++- .../service/openai_ws_http_bridge_test.go | 124 +++++++++++++++++- .../__tests__/AccountTestModal.spec.ts | 54 ++++++-- 15 files changed, 578 insertions(+), 72 deletions(-) create mode 100644 backend/internal/service/account_test_service_grok_test.go diff --git a/README.md b/README.md index 11d5511f7d..7f0d72e9db 100644 --- a/README.md +++ b/README.md @@ -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 diff --git a/backend/cmd/server/wire_gen.go b/backend/cmd/server/wire_gen.go index ee3ccf9489..12f78f2718 100644 --- a/backend/cmd/server/wire_gen.go +++ b/backend/cmd/server/wire_gen.go @@ -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) diff --git a/backend/internal/handler/admin/account_handler.go b/backend/internal/handler/admin/account_handler.go index 7967c15388..c43348528b 100644 --- a/backend/internal/handler/admin/account_handler.go +++ b/backend/internal/handler/admin/account_handler.go @@ -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() { diff --git a/backend/internal/handler/admin/account_handler_available_models_test.go b/backend/internal/handler/admin/account_handler_available_models_test.go index 0efbd6d434..eca929fc92 100644 --- a/backend/internal/handler/admin/account_handler_available_models_test.go +++ b/backend/internal/handler/admin/account_handler_available_models_test.go @@ -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(), diff --git a/backend/internal/handler/openai_gateway_handler.go b/backend/internal/handler/openai_gateway_handler.go index 1a7cf357dc..8602be7abd 100644 --- a/backend/internal/handler/openai_gateway_handler.go +++ b/backend/internal/handler/openai_gateway_handler.go @@ -1338,6 +1338,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 +1367,7 @@ func (h *OpenAIGatewayHandler) ResponsesWebSocket(c *gin.Context) { sessionHash, reqModel, failedAccountIDs, - service.OpenAIUpstreamTransportResponsesWebsocketV2, + requiredTransport, service.OpenAIEndpointCapabilityChatCompletions, false, requestPlatform, diff --git a/backend/internal/server/routes/gateway.go b/backend/internal/server/routes/gateway.go index 23b989e466..54fb5c4d68 100644 --- a/backend/internal/server/routes/gateway.go +++ b/backend/internal/server/routes/gateway.go @@ -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 } diff --git a/backend/internal/server/routes/gateway_test.go b/backend/internal/server/routes/gateway_test.go index 06038cd52e..0d8b6bd67a 100644 --- a/backend/internal/server/routes/gateway_test.go +++ b/backend/internal/server/routes/gateway_test.go @@ -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"}]}`)) diff --git a/backend/internal/service/account_test_service.go b/backend/internal/service/account_test_service.go index 7a9cd18a25..cc25943d27 100644 --- a/backend/internal/service/account_test_service.go +++ b/backend/internal/service/account_test_service.go @@ -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( diff --git a/backend/internal/service/account_test_service_grok_test.go b/backend/internal/service/account_test_service_grok_test.go new file mode 100644 index 0000000000..356fc62fb7 --- /dev/null +++ b/backend/internal/service/account_test_service_grok_test.go @@ -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"`) +} diff --git a/backend/internal/service/openai_gateway_grok_test.go b/backend/internal/service/openai_gateway_grok_test.go index 15dd6a3acb..168cddf514 100644 --- a/backend/internal/service/openai_gateway_grok_test.go +++ b/backend/internal/service/openai_gateway_grok_test.go @@ -271,6 +271,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 diff --git a/backend/internal/service/openai_gateway_messages.go b/backend/internal/service/openai_gateway_messages.go index 7d07fcc808..60839a0aa6 100644 --- a/backend/internal/service/openai_gateway_messages.go +++ b/backend/internal/service/openai_gateway_messages.go @@ -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) } } diff --git a/backend/internal/service/openai_ws_forwarder.go b/backend/internal/service/openai_ws_forwarder.go index a3e1eaf21e..efc805bc06 100644 --- a/backend/internal/service/openai_ws_forwarder.go +++ b/backend/internal/service/openai_ws_forwarder.go @@ -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, diff --git a/backend/internal/service/openai_ws_http_bridge.go b/backend/internal/service/openai_ws_http_bridge.go index 270964b580..0879d7e8fb 100644 --- a/backend/internal/service/openai_ws_http_bridge.go +++ b/backend/internal/service/openai_ws_http_bridge.go @@ -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 diff --git a/backend/internal/service/openai_ws_http_bridge_test.go b/backend/internal/service/openai_ws_http_bridge_test.go index 0a1d6b564b..0acb39ec83 100644 --- a/backend/internal/service/openai_ws_http_bridge_test.go +++ b/backend/internal/service/openai_ws_http_bridge_test.go @@ -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) diff --git a/frontend/src/components/admin/account/__tests__/AccountTestModal.spec.ts b/frontend/src/components/admin/account/__tests__/AccountTestModal.spec.ts index eb1a7b9d70..51001dadb1 100644 --- a/frontend/src/components/admin/account/__tests__/AccountTestModal.spec.ts +++ b/frontend/src/components/admin/account/__tests__/AccountTestModal.spec.ts @@ -59,17 +59,17 @@ function createStreamResponse(lines: string[]) { } as Response } -function mountModal() { +function mountModal(account: Record = { + 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: '' + }) + }) })