diff --git a/backend/internal/handler/openai_gateway_handler.go b/backend/internal/handler/openai_gateway_handler.go index 8602be7abd..2f3a95985d 100644 --- a/backend/internal/handler/openai_gateway_handler.go +++ b/backend/internal/handler/openai_gateway_handler.go @@ -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 diff --git a/backend/internal/handler/openai_gateway_handler_test.go b/backend/internal/handler/openai_gateway_handler_test.go index 5886388a9e..c4ccb5a024 100644 --- a/backend/internal/handler/openai_gateway_handler_test.go +++ b/backend/internal/handler/openai_gateway_handler_test.go @@ -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 diff --git a/backend/internal/service/openai_messages_dispatch.go b/backend/internal/service/openai_messages_dispatch.go index f2c1ad3c1b..f72b84a7f3 100644 --- a/backend/internal/service/openai_messages_dispatch.go +++ b/backend/internal/service/openai_messages_dispatch.go @@ -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 diff --git a/backend/internal/service/openai_messages_dispatch_test.go b/backend/internal/service/openai_messages_dispatch_test.go index a625aaddd4..e0b8ab0aa0 100644 --- a/backend/internal/service/openai_messages_dispatch_test.go +++ b/backend/internal/service/openai_messages_dispatch_test.go @@ -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")) +}