diff --git a/backend/internal/handler/admin/group_handler.go b/backend/internal/handler/admin/group_handler.go index 1c3ecb07e1..99833e912a 100644 --- a/backend/internal/handler/admin/group_handler.go +++ b/backend/internal/handler/admin/group_handler.go @@ -234,7 +234,7 @@ type UpdateGroupRequest struct { type CompositeRouteRequest struct { PublicModel string `json:"public_model" binding:"required"` MatchType string `json:"match_type" binding:"omitempty,oneof=exact prefix"` - TargetPlatform string `json:"target_platform" binding:"required,oneof=anthropic openai gemini antigravity grok"` + TargetPlatform string `json:"target_platform" binding:"required,oneof=anthropic openai gemini antigravity grok kimi zhipu deepseek"` UpstreamModel string `json:"upstream_model"` Endpoint string `json:"endpoint" binding:"omitempty,oneof=any messages count_tokens responses chat_completions embeddings images gemini"` Priority int `json:"priority"` diff --git a/backend/internal/handler/admin/group_handler_platform_test.go b/backend/internal/handler/admin/group_handler_platform_test.go index ca1703180c..aeba6d1c5a 100644 --- a/backend/internal/handler/admin/group_handler_platform_test.go +++ b/backend/internal/handler/admin/group_handler_platform_test.go @@ -71,13 +71,11 @@ func TestGroupPlatformBinding_RejectsInvalidPlatforms(t *testing.T) { } } -// 守住 composite 路由目标不放行 CN:CN 平台不可作为 composite 路由目标 -// (DetectModelPlatform/isConcreteRequestPlatform 均无 CN 分支,放行即打开半实现路径)。 -func TestCompositeRouteTargetPlatform_StillExcludesCNProviders(t *testing.T) { +func TestCompositeRouteTargetPlatform_AllowsCNProviders(t *testing.T) { for _, platform := range []string{"kimi", "zhipu", "deepseek"} { var req CompositeRouteRequest body := fmt.Sprintf(`{"public_model":"m","target_platform":%q}`, platform) - require.Error(t, bindGroupPlatformJSON(t, &req, body), - "composite target_platform %q 应保持被拒", platform) + require.NoError(t, bindGroupPlatformJSON(t, &req, body)) + require.Equal(t, platform, req.TargetPlatform) } } diff --git a/backend/internal/handler/composite_platform_test.go b/backend/internal/handler/composite_platform_test.go index 2fa79c538b..be1706fba3 100644 --- a/backend/internal/handler/composite_platform_test.go +++ b/backend/internal/handler/composite_platform_test.go @@ -22,18 +22,42 @@ func TestCompositeTargetPlatformAllowedResolvesKnownAllowedModel(t *testing.T) { require.Equal(t, service.PlatformOpenAI, platform) } -func TestOpenAICompatibleTextTargetAllowsCompositeGrokModel(t *testing.T) { +func TestOpenAICompatibleTextTargetAllowsCompositeProviders(t *testing.T) { gin.SetMode(gin.TestMode) - for _, path := range []string{"/v1/messages", "/v1/chat/completions"} { - c, _ := gin.CreateTestContext(httptest.NewRecorder()) - c.Request = httptest.NewRequest("POST", path, nil) - apiKey := &service.APIKey{Group: &service.Group{Platform: service.PlatformComposite}} + providers := []struct { + model string + platform string + }{ + {model: "grok-4.3", platform: service.PlatformGrok}, + {model: "kimi-k2-thinking", platform: service.PlatformKimi}, + {model: "glm-5.2", platform: service.PlatformZhipu}, + {model: "deepseek-v3.2", platform: service.PlatformDeepseek}, + } + for _, path := range []string{"/v1/messages", "/v1/chat/completions", "/v1/responses", "/v1/responses/input_tokens", "/v1/messages/count_tokens"} { + for _, provider := range providers { + c, _ := gin.CreateTestContext(httptest.NewRecorder()) + c.Request = httptest.NewRequest("POST", path, nil) + apiKey := &service.APIKey{Group: &service.Group{Platform: service.PlatformComposite}} - require.True(t, openAICompatibleTextTargetAllowed(c, apiKey, "grok-4.3"), "path=%s", path) - platform, ok := service.ResolvedTargetPlatformFromContext(c.Request.Context()) - require.True(t, ok, "path=%s", path) - require.Equal(t, service.PlatformGrok, platform, "path=%s", path) + require.True(t, openAICompatibleTextTargetAllowed(c, apiKey, provider.model), "path=%s model=%s", path, provider.model) + platform, ok := service.ResolvedTargetPlatformFromContext(c.Request.Context()) + require.True(t, ok, "path=%s model=%s", path, provider.model) + require.Equal(t, provider.platform, platform, "path=%s model=%s", path, provider.model) + } + } +} + +// WS ingress 对 CN 账号既过不了 transport 过滤、HTTP 桥也没有 Responses 转换, +// 放行只会把明确的策略拒绝换成 "no available account",因此 WS 白名单保持 openai+grok。 +func TestResponsesWebSocketCompositePlatformGuardKeepsOpenAIAndGrokOnly(t *testing.T) { + require.True(t, isResponsesWebSocketCompositePlatform(service.PlatformOpenAI)) + require.True(t, isResponsesWebSocketCompositePlatform(service.PlatformGrok)) + for _, platform := range []string{ + service.PlatformKimi, service.PlatformZhipu, service.PlatformDeepseek, + service.PlatformAnthropic, service.PlatformGemini, + } { + require.False(t, isResponsesWebSocketCompositePlatform(platform), "platform=%s", platform) } } diff --git a/backend/internal/handler/gateway_handler.go b/backend/internal/handler/gateway_handler.go index 07315d5efc..1128329be8 100644 --- a/backend/internal/handler/gateway_handler.go +++ b/backend/internal/handler/gateway_handler.go @@ -1147,10 +1147,12 @@ func (h *GatewayHandler) compositeAvailableModels(ctx context.Context, groupID * seen := make(map[string]struct{}) models := make([]string, 0) schedulablePlatforms := h.gatewayService.GetSchedulablePlatforms(ctx, groupID) - for _, platform := range []string{service.PlatformAnthropic, service.PlatformGemini, service.PlatformOpenAI, service.PlatformAntigravity, service.PlatformGrok} { + for _, platform := range []string{service.PlatformAnthropic, service.PlatformGemini, service.PlatformOpenAI, service.PlatformAntigravity, service.PlatformGrok, service.PlatformKimi, service.PlatformZhipu, service.PlatformDeepseek} { platformModels := h.gatewayService.GetAvailableModels(ctx, groupID, platform) if len(platformModels) == 0 { - if _, ok := schedulablePlatforms[platform]; ok { + // CN 供应商没有静态默认模型列表(defaultModelIDsForPlatform 的 + // default 分支是 Claude 列表),composite 下只暴露账号映射键。 + if _, ok := schedulablePlatforms[platform]; ok && !service.IsCNProvider(platform) { platformModels = defaultModelIDsForPlatform(platform) } } @@ -1372,7 +1374,7 @@ func defaultModelIDsForPlatform(platform string) []string { case service.PlatformComposite: ids := make([]string, 0) seen := make(map[string]struct{}) - for _, concretePlatform := range []string{service.PlatformAnthropic, service.PlatformGemini, service.PlatformOpenAI, service.PlatformAntigravity, service.PlatformGrok} { + for _, concretePlatform := range []string{service.PlatformAnthropic, service.PlatformGemini, service.PlatformOpenAI, service.PlatformAntigravity, service.PlatformGrok, service.PlatformKimi, service.PlatformZhipu, service.PlatformDeepseek} { for _, id := range defaultModelIDsForPlatform(concretePlatform) { if _, ok := seen[id]; ok { continue diff --git a/backend/internal/handler/gateway_models_test.go b/backend/internal/handler/gateway_models_test.go index 58f0190297..313b3216fb 100644 --- a/backend/internal/handler/gateway_models_test.go +++ b/backend/internal/handler/gateway_models_test.go @@ -7,6 +7,7 @@ import ( "net/http/httptest" "testing" + "github.com/Wei-Shaw/sub2api/internal/pkg/claude" middleware2 "github.com/Wei-Shaw/sub2api/internal/server/middleware" "github.com/Wei-Shaw/sub2api/internal/service" "github.com/gin-gonic/gin" @@ -321,6 +322,27 @@ func TestGatewayModels_CompositeCustomModelsListFiltersAcrossConcretePlatforms(t }, }, }, + { + ID: 4, + Platform: service.PlatformKimi, + Credentials: map[string]any{ + "model_mapping": map[string]any{"kimi-custom": "kimi-upstream"}, + }, + }, + { + ID: 5, + Platform: service.PlatformZhipu, + Credentials: map[string]any{ + "model_mapping": map[string]any{"glm-custom": "glm-upstream"}, + }, + }, + { + ID: 6, + Platform: service.PlatformDeepseek, + Credentials: map[string]any{ + "model_mapping": map[string]any{"deepseek-custom": "deepseek-upstream"}, + }, + }, }, }, }, @@ -335,7 +357,7 @@ func TestGatewayModels_CompositeCustomModelsListFiltersAcrossConcretePlatforms(t Platform: service.PlatformComposite, ModelsListConfig: service.GroupModelsListConfig{ Enabled: true, - Models: []string{"gemini-2.5-flash", "missing-model", "ag-custom-model", "gpt-5.5"}, + Models: []string{"gemini-2.5-flash", "missing-model", "ag-custom-model", "gpt-5.5", "kimi-custom", "glm-custom", "deepseek-custom"}, }, }, }) @@ -346,7 +368,7 @@ func TestGatewayModels_CompositeCustomModelsListFiltersAcrossConcretePlatforms(t var got gatewayModelsResponseForTest require.NoError(t, json.Unmarshal(rec.Body.Bytes(), &got)) - require.Equal(t, []string{"gemini-2.5-flash", "ag-custom-model", "gpt-5.5"}, modelIDsForTest(got.Data)) + require.Equal(t, []string{"gemini-2.5-flash", "ag-custom-model", "gpt-5.5", "kimi-custom", "glm-custom", "deepseek-custom"}, modelIDsForTest(got.Data)) } func TestGatewayModels_CompositeUnmappedAccountsFallbackToLinkedPlatformsOnly(t *testing.T) { @@ -385,6 +407,56 @@ func TestGatewayModels_CompositeUnmappedAccountsFallbackToLinkedPlatformsOnly(t require.NotContains(t, ids, "gemini-2.5-flash") } +// CN 供应商没有静态默认模型列表:composite 下无映射的可调度 CN 账号不得把 +// defaultModelIDsForPlatform default 分支的 Claude 列表挂到 CN 平台名下。 +func TestGatewayModels_CompositeUnmappedCNAccountsContributeNoDefaults(t *testing.T) { + gin.SetMode(gin.TestMode) + + groupID := int64(35) + h := newGatewayModelsHandlerForTest( + &gatewayModelsAccountRepoStub{ + byGroup: map[int64][]service.Account{ + groupID: { + {ID: 1, Platform: service.PlatformOpenAI}, + {ID: 2, Platform: service.PlatformKimi}, + {ID: 3, Platform: service.PlatformZhipu}, + {ID: 4, Platform: service.PlatformDeepseek}, + }, + }, + }, + ) + + rec := httptest.NewRecorder() + c, _ := gin.CreateTestContext(rec) + c.Request = httptest.NewRequest(http.MethodGet, "/v1/models", nil) + c.Set(string(middleware2.ContextKeyAPIKey), &service.APIKey{ + Group: &service.Group{ID: groupID, Platform: service.PlatformComposite}, + }) + + h.Models(c) + + require.Equal(t, http.StatusOK, rec.Code) + + var got gatewayModelsResponseForTest + require.NoError(t, json.Unmarshal(rec.Body.Bytes(), &got)) + + ids := modelIDsForTest(got.Data) + require.Contains(t, ids, "gpt-5.5") + require.NotContains(t, ids, "claude-sonnet-4-6") +} + +// 独立 CN 分组沿用 default 分支的 Claude 默认列表(Claude Code 客户端请求的 +// 就是这些模型名并经账号 model_mapping 转换),composite 支持不得改变该回退。 +func TestDefaultModelIDsForPlatform_CNProvidersKeepClaudeDefaults(t *testing.T) { + want := make([]string, 0, len(claude.DefaultModels)) + for _, model := range claude.DefaultModels { + want = append(want, model.ID) + } + for _, platform := range []string{service.PlatformKimi, service.PlatformZhipu, service.PlatformDeepseek} { + require.Equal(t, want, defaultModelIDsForPlatform(platform), "platform=%s", platform) + } +} + func TestGatewayModels_CustomModelsListKeepsConcreteModelAllowedByWildcardMapping(t *testing.T) { gin.SetMode(gin.TestMode) diff --git a/backend/internal/handler/openai_gateway_cn_dispatch_test.go b/backend/internal/handler/openai_gateway_cn_dispatch_test.go index b7f11c5462..c49201bd12 100644 --- a/backend/internal/handler/openai_gateway_cn_dispatch_test.go +++ b/backend/internal/handler/openai_gateway_cn_dispatch_test.go @@ -3,27 +3,74 @@ package handler // CN 分组 /v1/messages 调度闸门回归(修复:正常途径创建的 CN 分组曾恒 403): // sanitizeGroupMessagesDispatchFields 对非 openai 平台强制 AllowMessagesDispatch // =false,故 CN 分组必须与 grok 一样在闸门处豁免,否则原生 Anthropic 直通 -//(Claude Code 主用例)永远不可达。 +//(Claude Code 主用例)永远不可达。composite 分组同理:sanitize 对 composite +// 恒置 false,解析到 grok/CN 目标时必须按目标平台豁免,解析到 openai 目标 +// 仍受开关控制。 import ( + "net/http/httptest" "testing" "github.com/Wei-Shaw/sub2api/internal/service" + "github.com/gin-gonic/gin" "github.com/stretchr/testify/require" ) func TestAllowOpenAICompatibleMessagesDispatch_CNProvidersExempt(t *testing.T) { - require.True(t, allowOpenAICompatibleMessagesDispatch(nil), "无 key 保持放行") + require.True(t, allowOpenAICompatibleMessagesDispatch(nil, nil), "无 key 保持放行") for _, platform := range []string{service.PlatformKimi, service.PlatformZhipu, service.PlatformDeepseek, service.PlatformGrok} { apiKey := &service.APIKey{Group: &service.Group{Platform: platform, AllowMessagesDispatch: false}} - require.True(t, allowOpenAICompatibleMessagesDispatch(apiKey), + require.True(t, allowOpenAICompatibleMessagesDispatch(nil, apiKey), "%s 分组必须豁免 allow_messages_dispatch 闸门", platform) } // 非回归:openai 分组仍受开关控制。 openaiOff := &service.APIKey{Group: &service.Group{Platform: service.PlatformOpenAI, AllowMessagesDispatch: false}} - require.False(t, allowOpenAICompatibleMessagesDispatch(openaiOff)) + require.False(t, allowOpenAICompatibleMessagesDispatch(nil, openaiOff)) openaiOn := &service.APIKey{Group: &service.Group{Platform: service.PlatformOpenAI, AllowMessagesDispatch: true}} - require.True(t, allowOpenAICompatibleMessagesDispatch(openaiOn)) + require.True(t, allowOpenAICompatibleMessagesDispatch(nil, openaiOn)) +} + +func TestAllowOpenAICompatibleMessagesDispatch_CompositeResolvedTargets(t *testing.T) { + gin.SetMode(gin.TestMode) + + newCompositeCtx := func(model string) (*gin.Context, *service.APIKey) { + c, _ := gin.CreateTestContext(httptest.NewRecorder()) + c.Request = httptest.NewRequest("POST", "/v1/messages", nil) + apiKey := &service.APIKey{Group: &service.Group{Platform: service.PlatformComposite, AllowMessagesDispatch: false}} + ensureCompositeTargetPlatform(c, apiKey, model) + return c, apiKey + } + + // 解析到 grok/CN 目标:与对应独立分组同语义豁免。 + for _, model := range []string{"grok-4.3", "kimi-k2-thinking", "glm-5.2", "deepseek-v3.2"} { + c, apiKey := newCompositeCtx(model) + require.True(t, allowOpenAICompatibleMessagesDispatch(c, apiKey), "model=%s", model) + } + + // 解析到 openai 目标:仍受开关控制(composite 被 sanitize 恒置 false ⇒ 拒绝)。 + c, apiKey := newCompositeCtx("gpt-5.5") + require.False(t, allowOpenAICompatibleMessagesDispatch(c, apiKey)) + + // 未解析出目标平台:保持拒绝,不放宽。 + cNone, _ := gin.CreateTestContext(httptest.NewRecorder()) + cNone.Request = httptest.NewRequest("POST", "/v1/messages", nil) + require.False(t, allowOpenAICompatibleMessagesDispatch(cNone, + &service.APIKey{Group: &service.Group{Platform: service.PlatformComposite, AllowMessagesDispatch: false}})) +} + +// composite 解析到 grok/CN 目标时,Group 级调度映射(gpt-5.x 默认值为 openai +// 专属)不得注入,模型改写完全交给账号级 model_mapping。 +func TestResolveOpenAIMessagesDispatchMappedModel_CompositeCNTargetsSkipGroupMapping(t *testing.T) { + gin.SetMode(gin.TestMode) + + for _, model := range []string{"kimi-k2-thinking", "glm-5.2", "deepseek-v3.2", "grok-4.3"} { + c, _ := gin.CreateTestContext(httptest.NewRecorder()) + c.Request = httptest.NewRequest("POST", "/v1/messages", nil) + apiKey := &service.APIKey{Group: &service.Group{Platform: service.PlatformComposite}} + ensureCompositeTargetPlatform(c, apiKey, model) + + require.Empty(t, resolveOpenAIMessagesDispatchMappedModel(c, apiKey, "claude-sonnet-4-5-20250929"), "model=%s", model) + } } diff --git a/backend/internal/handler/openai_gateway_count_tokens.go b/backend/internal/handler/openai_gateway_count_tokens.go index f300b31b95..ee6e86c4f7 100644 --- a/backend/internal/handler/openai_gateway_count_tokens.go +++ b/backend/internal/handler/openai_gateway_count_tokens.go @@ -60,7 +60,7 @@ func (h *OpenAIGatewayHandler) ResponsesInputTokens(c *gin.Context) { } reqModel := strings.TrimSpace(modelResult.String()) ensureCompositeTargetPlatform(c, apiKey, reqModel) - if !compositeTargetPlatformAllowed(c, apiKey, reqModel, service.PlatformOpenAI, service.PlatformGrok) { + if !openAICompatibleTextTargetAllowed(c, apiKey, reqModel) { h.errorResponse(c, http.StatusBadRequest, "invalid_request_error", "Model is not supported by this OpenAI-compatible endpoint for composite groups") return } @@ -204,7 +204,7 @@ func (h *OpenAIGatewayHandler) CountTokens(c *gin.Context) { zap.Any("group_id", apiKey.GroupID), ) - if !allowOpenAICompatibleMessagesDispatch(apiKey) { + if !allowOpenAICompatibleMessagesDispatch(c, apiKey) { h.anthropicErrorResponse(c, http.StatusForbidden, "permission_error", "This group does not allow /v1/messages dispatch") return @@ -243,12 +243,14 @@ func (h *OpenAIGatewayHandler) CountTokens(c *gin.Context) { reqModel := parsedReq.Model ensureCompositeTargetPlatform(c, apiKey, reqModel) - if !compositeTargetPlatformAllowed(c, apiKey, reqModel, service.PlatformOpenAI) { + // composite+grok 在路由层已分流到 GrokCountTokens,这里可达的目标平台是 + // openai 与 CN 供应商;CN 账号由 ForwardCountTokensAsAnthropic 本地估算。 + if !openAICompatibleTextTargetAllowed(c, apiKey, reqModel) { h.anthropicErrorResponse(c, http.StatusBadRequest, "invalid_request_error", "Model is not supported by this OpenAI-compatible endpoint for composite groups") return } routingModel := service.NormalizeOpenAICompatRequestedModel(reqModel) - preferredMappedModel := resolveOpenAIMessagesDispatchMappedModel(apiKey, reqModel) + preferredMappedModel := resolveOpenAIMessagesDispatchMappedModel(c, apiKey, reqModel) reqLog = reqLog.With(zap.String("model", reqModel), zap.Bool("stream", parsedReq.Stream)) setOpsRequestContext(c, reqModel, false) diff --git a/backend/internal/handler/openai_gateway_handler.go b/backend/internal/handler/openai_gateway_handler.go index f632803b7d..84b6417e5c 100644 --- a/backend/internal/handler/openai_gateway_handler.go +++ b/backend/internal/handler/openai_gateway_handler.go @@ -98,10 +98,18 @@ func openAIForwardSucceededForScheduling(result *service.OpenAIForwardResult) bo return result.SucceededForScheduling() } -func resolveOpenAIMessagesDispatchMappedModel(apiKey *service.APIKey, requestedModel string) string { +func resolveOpenAIMessagesDispatchMappedModel(c *gin.Context, apiKey *service.APIKey, requestedModel string) string { if apiKey == nil || apiKey.Group == nil { return "" } + // composite 解析到 grok/CN 目标时调度级映射不适用(Group 级映射的 gpt-5.x + // 默认值是 openai 专属,发给这些上游必错),模型改写交给账号级 model_mapping。 + if apiKey.Group.Platform == service.PlatformComposite && c != nil && c.Request != nil { + if platform, ok := service.ResolvedTargetPlatformFromContext(c.Request.Context()); ok && + (platform == service.PlatformGrok || service.IsCNProvider(platform)) { + return "" + } + } return strings.TrimSpace(apiKey.Group.ResolveMessagesDispatchModel(requestedModel)) } @@ -190,7 +198,7 @@ func openAIResponsesRequiredCapabilityForRequest(imageIntent bool, needsResponse return openAIResponsesRequiredCapability(imageIntent, platform) } -func allowOpenAICompatibleMessagesDispatch(apiKey *service.APIKey) bool { +func allowOpenAICompatibleMessagesDispatch(c *gin.Context, apiKey *service.APIKey) bool { if apiKey == nil || apiKey.Group == nil { return true } @@ -204,6 +212,15 @@ func allowOpenAICompatibleMessagesDispatch(apiKey *service.APIKey) bool { if service.IsCNProvider(apiKey.Group.Platform) { return true } + // composite 分组解析到 grok/CN 目标时与对应独立分组同语义豁免:sanitize + // 对 composite 同样恒置 false,不豁免则这些目标的 /v1/messages 永远 403; + // 解析到 openai 目标仍受开关控制,维持现状。 + if apiKey.Group.Platform == service.PlatformComposite && c != nil && c.Request != nil { + if platform, ok := service.ResolvedTargetPlatformFromContext(c.Request.Context()); ok && + (platform == service.PlatformGrok || service.IsCNProvider(platform)) { + return true + } + } return apiKey.Group.AllowMessagesDispatch } @@ -213,6 +230,19 @@ func openAICompatibleTextTargetAllowed(c *gin.Context, apiKey *service.APIKey, m service.PlatformKimi, service.PlatformZhipu, service.PlatformDeepseek) } +// isResponsesWebSocketCompositePlatform 限定 composite 分组在 Responses WebSocket +// 上可服务的目标平台。CN 供应商(kimi/zhipu/deepseek)刻意排除:其账号无法通过 +// WSv2 ingress 的 transport 过滤,且 WS HTTP 桥没有面向 CN 的 Responses 转换, +// 放行只会把明确的策略拒绝变成误导性的 "no available account"。 +func isResponsesWebSocketCompositePlatform(platform string) bool { + switch platform { + case service.PlatformOpenAI, service.PlatformGrok: + return true + default: + return false + } +} + // NewOpenAIGatewayHandler creates a new OpenAIGatewayHandler func NewOpenAIGatewayHandler( gatewayService *service.OpenAIGatewayService, @@ -333,7 +363,7 @@ func (h *OpenAIGatewayHandler) Responses(c *gin.Context) { } reqModel := modelResult.String() ensureCompositeTargetPlatform(c, apiKey, reqModel) - if !compositeTargetPlatformAllowed(c, apiKey, reqModel, service.PlatformOpenAI, service.PlatformGrok) { + if !openAICompatibleTextTargetAllowed(c, apiKey, reqModel) { h.errorResponse(c, http.StatusBadRequest, "invalid_request_error", "Model is not supported by this OpenAI-compatible endpoint for composite groups") return } @@ -944,7 +974,7 @@ func (h *OpenAIGatewayHandler) Messages(c *gin.Context) { ) // 检查分组是否允许 /v1/messages 调度 - if !allowOpenAICompatibleMessagesDispatch(apiKey) { + if !allowOpenAICompatibleMessagesDispatch(c, apiKey) { h.anthropicErrorResponse(c, http.StatusForbidden, "permission_error", "This group does not allow /v1/messages dispatch") return @@ -987,7 +1017,7 @@ func (h *OpenAIGatewayHandler) Messages(c *gin.Context) { } bindOpenAIReasoningEffortPolicyForMessagesRequest(c, apiKey, body) routingModel := service.NormalizeOpenAICompatRequestedModel(reqModel) - preferredMappedModel := resolveOpenAIMessagesDispatchMappedModel(apiKey, reqModel) + preferredMappedModel := resolveOpenAIMessagesDispatchMappedModel(c, apiKey, reqModel) reqStream := gjson.GetBytes(body, "stream").Bool() reqLog = reqLog.With(zap.String("model", reqModel), zap.Bool("stream", reqStream)) @@ -1750,7 +1780,7 @@ func (h *OpenAIGatewayHandler) ResponsesWebSocket(c *gin.Context) { ctx = c.Request.Context() if apiKey.Group != nil && apiKey.Group.Platform == service.PlatformComposite { platform, ok := service.ResolvedTargetPlatformFromContext(ctx) - if !ok || (platform != service.PlatformOpenAI && platform != service.PlatformGrok) { + if !ok || !isResponsesWebSocketCompositePlatform(platform) { closeOpenAIClientWS(wsConn, coderws.StatusPolicyViolation, "Responses WebSocket API only supports OpenAI-compatible models for composite groups") return } diff --git a/backend/internal/handler/openai_gateway_handler_test.go b/backend/internal/handler/openai_gateway_handler_test.go index 6e8a687d77..898449cc29 100644 --- a/backend/internal/handler/openai_gateway_handler_test.go +++ b/backend/internal/handler/openai_gateway_handler_test.go @@ -646,21 +646,21 @@ func TestResolveOpenAIMessagesDispatchMappedModel(t *testing.T) { }, }, } - require.Equal(t, "gpt-5.4-mini", resolveOpenAIMessagesDispatchMappedModel(apiKey, "claude-sonnet-4-5-20250929")) - require.Equal(t, "gpt-5.6-sol", resolveOpenAIMessagesDispatchMappedModel(apiKey, "claude-fable-5")) + require.Equal(t, "gpt-5.4-mini", resolveOpenAIMessagesDispatchMappedModel(nil, apiKey, "claude-sonnet-4-5-20250929")) + require.Equal(t, "gpt-5.6-sol", resolveOpenAIMessagesDispatchMappedModel(nil, apiKey, "claude-fable-5")) }) t.Run("uses_family_default_when_no_override", func(t *testing.T) { apiKey := &service.APIKey{Group: &service.Group{}} - require.Equal(t, "gpt-5.4", resolveOpenAIMessagesDispatchMappedModel(apiKey, "claude-opus-4-6")) - require.Equal(t, "gpt-5.3-codex", resolveOpenAIMessagesDispatchMappedModel(apiKey, "claude-sonnet-4-5-20250929")) - require.Equal(t, "gpt-5.4-mini", resolveOpenAIMessagesDispatchMappedModel(apiKey, "claude-haiku-4-5-20251001")) + require.Equal(t, "gpt-5.4", resolveOpenAIMessagesDispatchMappedModel(nil, apiKey, "claude-opus-4-6")) + require.Equal(t, "gpt-5.3-codex", resolveOpenAIMessagesDispatchMappedModel(nil, apiKey, "claude-sonnet-4-5-20250929")) + require.Equal(t, "gpt-5.4-mini", resolveOpenAIMessagesDispatchMappedModel(nil, apiKey, "claude-haiku-4-5-20251001")) }) t.Run("returns_empty_for_non_claude_or_missing_group", func(t *testing.T) { - require.Empty(t, resolveOpenAIMessagesDispatchMappedModel(nil, "claude-sonnet-4-5-20250929")) - require.Empty(t, resolveOpenAIMessagesDispatchMappedModel(&service.APIKey{}, "claude-sonnet-4-5-20250929")) - require.Empty(t, resolveOpenAIMessagesDispatchMappedModel(&service.APIKey{Group: &service.Group{}}, "gpt-5.4")) + require.Empty(t, resolveOpenAIMessagesDispatchMappedModel(nil, nil, "claude-sonnet-4-5-20250929")) + require.Empty(t, resolveOpenAIMessagesDispatchMappedModel(nil, &service.APIKey{}, "claude-sonnet-4-5-20250929")) + require.Empty(t, resolveOpenAIMessagesDispatchMappedModel(nil, &service.APIKey{Group: &service.Group{}}, "gpt-5.4")) }) t.Run("grok_group_maps_claude_cli_model_to_grok_default", func(t *testing.T) { @@ -672,8 +672,8 @@ func TestResolveOpenAIMessagesDispatchMappedModel(t *testing.T) { Platform: service.PlatformGrok, }, } - require.Equal(t, "grok-4.5", resolveOpenAIMessagesDispatchMappedModel(apiKey, "claude-sonnet-4-5")) - require.Empty(t, resolveOpenAIMessagesDispatchMappedModel(apiKey, "grok")) + require.Equal(t, "grok-4.5", resolveOpenAIMessagesDispatchMappedModel(nil, apiKey, "claude-sonnet-4-5")) + require.Empty(t, resolveOpenAIMessagesDispatchMappedModel(nil, apiKey, "grok")) }) t.Run("does_not_fall_back_to_group_default_mapped_model", func(t *testing.T) { @@ -682,8 +682,8 @@ func TestResolveOpenAIMessagesDispatchMappedModel(t *testing.T) { DefaultMappedModel: "gpt-5.4", }, } - require.Empty(t, resolveOpenAIMessagesDispatchMappedModel(apiKey, "gpt-5.4")) - require.Equal(t, "gpt-5.3-codex", resolveOpenAIMessagesDispatchMappedModel(apiKey, "claude-sonnet-4-5-20250929")) + require.Empty(t, resolveOpenAIMessagesDispatchMappedModel(nil, apiKey, "gpt-5.4")) + require.Equal(t, "gpt-5.3-codex", resolveOpenAIMessagesDispatchMappedModel(nil, apiKey, "claude-sonnet-4-5-20250929")) }) } diff --git a/backend/internal/service/admin_group.go b/backend/internal/service/admin_group.go index 83467c9849..7b64e4b60f 100644 --- a/backend/internal/service/admin_group.go +++ b/backend/internal/service/admin_group.go @@ -267,7 +267,7 @@ func defaultAllowImageGenerationForPlatform(platform string) bool { func compositeDefaultModelsListCandidateIDs() []string { seen := make(map[string]struct{}) ids := make([]string, 0) - for _, platform := range []string{PlatformAnthropic, PlatformGemini, PlatformOpenAI, PlatformAntigravity, PlatformGrok} { + for _, platform := range []string{PlatformAnthropic, PlatformGemini, PlatformOpenAI, PlatformAntigravity, PlatformGrok, PlatformKimi, PlatformZhipu, PlatformDeepseek} { for _, id := range defaultModelsListCandidateIDs(platform) { if _, ok := seen[id]; ok { continue diff --git a/backend/internal/service/admin_service_composite_group_test.go b/backend/internal/service/admin_service_composite_group_test.go index 5530720f5e..f1c00482b1 100644 --- a/backend/internal/service/admin_service_composite_group_test.go +++ b/backend/internal/service/admin_service_composite_group_test.go @@ -6,6 +6,7 @@ import ( "context" "testing" + "github.com/Wei-Shaw/sub2api/internal/pkg/claude" "github.com/stretchr/testify/require" ) @@ -174,6 +175,13 @@ func TestAdminService_CompositeModelsListCandidatesIncludeConcreteAccountMapping "model_mapping": map[string]any{"gemini-custom": "gemini-2.5-flash"}, }, }, + { + ID: 3, + Platform: PlatformKimi, + Credentials: map[string]any{ + "model_mapping": map[string]any{"kimi-custom": "kimi-k2"}, + }, + }, }, } groupRepo := &groupRepoStubForAdmin{ @@ -188,6 +196,19 @@ func TestAdminService_CompositeModelsListCandidatesIncludeConcreteAccountMapping require.NoError(t, err) require.Contains(t, candidates, "gpt-custom") require.Contains(t, candidates, "gemini-custom") + require.Contains(t, candidates, "kimi-custom") require.Contains(t, candidates, "gpt-5.5") require.Contains(t, candidates, "gemini-2.5-flash") } + +// 独立 CN 分组的模型列表候选沿用 default 分支的 Claude 默认列表; +// composite 支持不得改变独立分组的候选语义。 +func TestAdminService_CNProviderModelsListCandidatesKeepClaudeDefaults(t *testing.T) { + want := make([]string, 0, len(claude.DefaultModels)) + for _, model := range claude.DefaultModels { + want = append(want, model.ID) + } + for _, platform := range []string{PlatformKimi, PlatformZhipu, PlatformDeepseek} { + require.Equal(t, want, defaultModelsListCandidateIDs(platform), "platform=%s", platform) + } +} diff --git a/backend/internal/service/channel_service.go b/backend/internal/service/channel_service.go index a2c4bf6ebf..ee732dbeef 100644 --- a/backend/internal/service/channel_service.go +++ b/backend/internal/service/channel_service.go @@ -357,7 +357,7 @@ func isPlatformPricingMatch(groupPlatform, pricingPlatform string) bool { // fallback used before a request target has been resolved. func matchingPlatforms(groupPlatform string) []string { if groupPlatform == PlatformComposite { - return []string{PlatformAnthropic, PlatformGemini, PlatformOpenAI, PlatformAntigravity, PlatformGrok} + return []string{PlatformAnthropic, PlatformGemini, PlatformOpenAI, PlatformAntigravity, PlatformGrok, PlatformKimi, PlatformZhipu, PlatformDeepseek} } return []string{groupPlatform} } diff --git a/backend/internal/service/channel_service_test.go b/backend/internal/service/channel_service_test.go index 9a81913b04..e443b80b59 100644 --- a/backend/internal/service/channel_service_test.go +++ b/backend/internal/service/channel_service_test.go @@ -2041,6 +2041,9 @@ func TestIsPlatformPricingMatch(t *testing.T) { {"gemini does NOT match anthropic", PlatformGemini, PlatformAnthropic, false}, {"composite matches openai pricing", PlatformComposite, PlatformOpenAI, true}, {"composite matches gemini pricing", PlatformComposite, PlatformGemini, true}, + {"composite matches kimi pricing", PlatformComposite, PlatformKimi, true}, + {"composite matches zhipu pricing", PlatformComposite, PlatformZhipu, true}, + {"composite matches deepseek pricing", PlatformComposite, PlatformDeepseek, true}, {"empty string matches nothing", "", PlatformAnthropic, false}, {"empty string matches empty", "", "", true}, } @@ -2066,7 +2069,7 @@ func TestMatchingPlatforms(t *testing.T) { {"anthropic returns itself", PlatformAnthropic, []string{PlatformAnthropic}}, {"gemini returns itself", PlatformGemini, []string{PlatformGemini}}, {"openai returns itself", PlatformOpenAI, []string{PlatformOpenAI}}, - {"composite returns concrete platforms", PlatformComposite, []string{PlatformAnthropic, PlatformGemini, PlatformOpenAI, PlatformAntigravity, PlatformGrok}}, + {"composite returns concrete platforms", PlatformComposite, []string{PlatformAnthropic, PlatformGemini, PlatformOpenAI, PlatformAntigravity, PlatformGrok, PlatformKimi, PlatformZhipu, PlatformDeepseek}}, } for _, tt := range tests { diff --git a/backend/internal/service/composite_platform.go b/backend/internal/service/composite_platform.go index b5def4ff54..e4d836febe 100644 --- a/backend/internal/service/composite_platform.go +++ b/backend/internal/service/composite_platform.go @@ -106,6 +106,12 @@ func DetectModelPlatform(model string) (string, bool) { return PlatformGemini, true case "xai", "x-ai", "grok": return PlatformGrok, true + case "kimi", "moonshot": + return PlatformKimi, true + case "zhipu", "glm", "bigmodel": + return PlatformZhipu, true + case "deepseek": + return PlatformDeepseek, true } if rest != "" { normalized = strings.TrimPrefix(rest, "models/") @@ -133,6 +139,13 @@ func DetectModelPlatform(model string) (string, bool) { return PlatformGemini, true case normalized == "grok" || strings.HasPrefix(normalized, "grok-"): return PlatformGrok, true + case strings.HasPrefix(normalized, "kimi-"), + strings.HasPrefix(normalized, "moonshot-"): + return PlatformKimi, true + case strings.HasPrefix(normalized, "glm-"): + return PlatformZhipu, true + case strings.HasPrefix(normalized, "deepseek-"): + return PlatformDeepseek, true default: return "", false } @@ -179,7 +192,8 @@ func (s *GatewayService) resolveCompositeRouteDecision(ctx context.Context, grou func isConcreteRequestPlatform(platform string) bool { switch platform { - case PlatformAnthropic, PlatformOpenAI, PlatformGemini, PlatformAntigravity, PlatformGrok: + case PlatformAnthropic, PlatformOpenAI, PlatformGemini, PlatformAntigravity, PlatformGrok, + PlatformKimi, PlatformZhipu, PlatformDeepseek: return true default: return false diff --git a/backend/internal/service/composite_platform_test.go b/backend/internal/service/composite_platform_test.go index 290fc4ac0e..e145634889 100644 --- a/backend/internal/service/composite_platform_test.go +++ b/backend/internal/service/composite_platform_test.go @@ -25,6 +25,10 @@ func TestDetectModelPlatform(t *testing.T) { {name: "learnlm", model: "learnlm-2.0-flash-experimental", platform: PlatformGemini, ok: true}, {name: "grok", model: "grok-4", platform: PlatformGrok, ok: true}, {name: "xai prefix", model: "xai/grok-4", platform: PlatformGrok, ok: true}, + {name: "kimi", model: "kimi-k2-thinking", platform: PlatformKimi, ok: true}, + {name: "moonshot prefix", model: "moonshot/moonshot-v1-32k", platform: PlatformKimi, ok: true}, + {name: "zhipu", model: "glm-5.2", platform: PlatformZhipu, ok: true}, + {name: "deepseek", model: "deepseek-v4-pro", platform: PlatformDeepseek, ok: true}, {name: "unknown", model: "llama-4-maverick", ok: false}, } @@ -59,7 +63,14 @@ func TestCompositeGroupSchedulerHasAllCanonicalPlatformBuckets(t *testing.T) { platforms = append(platforms, platform) } require.ElementsMatch(t, - []string{PlatformAnthropic, PlatformGemini, PlatformOpenAI, PlatformAntigravity, PlatformGrok}, + []string{PlatformAnthropic, PlatformGemini, PlatformOpenAI, PlatformAntigravity, PlatformGrok, PlatformKimi, PlatformZhipu, PlatformDeepseek}, platforms, ) } + +func TestCompositeConcretePlatformsIncludeCNProviders(t *testing.T) { + for _, platform := range []string{PlatformKimi, PlatformZhipu, PlatformDeepseek} { + require.True(t, isConcreteRequestPlatform(platform)) + require.True(t, canCopyAccountsFromGroupPlatform(PlatformComposite, platform)) + } +} diff --git a/backend/internal/service/scheduler_snapshot_bulk_event_test.go b/backend/internal/service/scheduler_snapshot_bulk_event_test.go index 00454a498d..0f6003a01a 100644 --- a/backend/internal/service/scheduler_snapshot_bulk_event_test.go +++ b/backend/internal/service/scheduler_snapshot_bulk_event_test.go @@ -114,6 +114,21 @@ func TestSchedulerBulkAccountEventScopesOpenAIRebuildToFreshPlatform(t *testing. require.Empty(t, deleted) } +func TestSchedulerBulkAccountEventScopesCNRebuildToFreshPlatform(t *testing.T) { + for _, platform := range []string{PlatformKimi, PlatformZhipu, PlatformDeepseek} { + t.Run(platform, func(t *testing.T) { + cache := newBulkEventSnapshotCache() + repo := newBulkEventAccountRepo(&Account{ID: 1, Platform: platform, GroupIDs: []int64{12}}) + svc := newBulkEventTestService(cache, repo) + + err := svc.handleBulkAccountEvent(context.Background(), bulkEventPayload([]int64{1}, []int64{11}), make(map[batchSeenKey]struct{})) + + require.NoError(t, err) + require.ElementsMatch(t, schedulerBucketsForTest([]int64{11, 12}, platform), cache.capturedBuckets()) + }) + } +} + func TestSchedulerBulkAccountEventRebuildsOpenAIUngroupedBucket(t *testing.T) { cache := newBulkEventSnapshotCache() repo := newBulkEventAccountRepo(&Account{ID: 6, Platform: PlatformOpenAI}) diff --git a/backend/internal/service/scheduler_snapshot_full_rebuild_lifecycle_test.go b/backend/internal/service/scheduler_snapshot_full_rebuild_lifecycle_test.go index fa918a0862..e52adbb4c8 100644 --- a/backend/internal/service/scheduler_snapshot_full_rebuild_lifecycle_test.go +++ b/backend/internal/service/scheduler_snapshot_full_rebuild_lifecycle_test.go @@ -290,13 +290,13 @@ func TestSchedulerFullRebuildActiveTombstoneDoesNotBlockFollowingGroupEvent(t *t require.Equal(t, 1, activeCalls) require.Zero(t, fallbackCalls) require.Equal(t, []int64{groupID, groupID}, freshCalls) - require.Len(t, cache.tokens(), 24, "full rebuild and the following group event must each run fresh authority") + require.Len(t, cache.tokens(), 36, "full rebuild and the following group event must each run fresh authority") _, reopenHeld := cache.lifecycleMutationLeaseStates() - require.Len(t, reopenHeld, 24) + require.Len(t, reopenHeld, 36) for _, held := range reopenHeld { require.True(t, held) } - require.Equal(t, 21, accounts.callCount()) + require.Equal(t, 30, accounts.callCount()) } func TestSchedulerFullRebuildGlobalReadErrorsFailBeforeMutationOrDB(t *testing.T) { @@ -367,15 +367,15 @@ func TestSchedulerFullRebuildFreshActivePreparesEveryTokenBeforeFirstDB(t *testi capturesAtFirstDB = cache.captureAttemptCount() held, reopenCount := cache.leaseHeldAndTokenCount() require.False(t, held) - require.Equal(t, 12, reopenCount) - require.Equal(t, 13, capturesAtFirstDB, "C(0) and the historical bucket must be captured before DB") + require.Equal(t, 18, reopenCount) + require.Equal(t, 19, capturesAtFirstDB, "C(0) and the historical bucket must be captured before DB") } svc := newFullRebuildLifecycleService(cache, nil, accounts, groups, config.RunModeStandard) require.NoError(t, svc.rebuildFullSnapshot(context.Background(), "test")) require.Equal(t, capturesAtFirstDB, cache.captureAttemptCount()) - require.Equal(t, 15, accounts.callCount()) - require.Equal(t, 8, accounts.groupCallCount(groupID)) + require.Equal(t, 21, accounts.callCount()) + require.Equal(t, 11, accounts.groupCallCount(groupID)) _, historicalPublished := cache.counts(historical) require.Equal(t, 1, historicalPublished) activeCalls, fallbackCalls, freshCalls := groups.stats() @@ -398,7 +398,7 @@ func TestSchedulerFullRebuildOrdinaryCaptureErrorReturnsBeforeFirstDB(t *testing err := svc.rebuildFullSnapshot(context.Background(), "test") require.ErrorIs(t, err, wantErr) - require.Equal(t, 14, cache.captureAttemptCount(), "all canonical and ordinary captures must be attempted before returning") + require.Equal(t, 20, cache.captureAttemptCount(), "all canonical and ordinary captures must be attempted before returning") require.Zero(t, accounts.callCount()) require.Zero(t, cache.totalSetAttempts()) } @@ -416,8 +416,8 @@ func TestSchedulerFullRebuildPreservesGroupZeroActiveHistoricalAndInvalidRegistr svc := newFullRebuildLifecycleService(cache, nil, accounts, groups, config.RunModeStandard) require.NoError(t, svc.rebuildFullSnapshot(context.Background(), "test")) - require.Equal(t, 27, cache.captureAttemptCount()) - require.Equal(t, 17, accounts.callCount()) + require.Equal(t, 39, cache.captureAttemptCount()) + require.Equal(t, 23, accounts.callCount()) groups.mu.Lock() require.Equal(t, 1, groups.listCalls) groups.mu.Unlock() @@ -456,7 +456,7 @@ func TestSchedulerFullRebuildActiveTombstoneFreshInactiveOrMissingFiltersAllGrou require.NoError(t, svc.rebuildFullSnapshot(context.Background(), "test")) require.Zero(t, accounts.groupCallCount(groupID)) - require.Equal(t, 7, accounts.groupCallCount(0)) + require.Equal(t, 10, accounts.groupCallCount(0)) require.Empty(t, cache.tokens()) require.Equal(t, bucketStrings(append(canonical, historical)), bucketStrings(cache.retiredBuckets())) for _, bucket := range append(canonical, historical) { @@ -528,7 +528,7 @@ func TestSchedulerFullRebuildPartialLifecycleFailureReturnsBeforeDBAndRetries(t require.Equal(t, []int64{1, 2}, freshCalls) require.Zero(t, accounts.callCount()) require.Zero(t, cache.totalSetAttempts()) - require.Equal(t, 13, len(cache.retiredBuckets())) + require.Equal(t, 19, len(cache.retiredBuckets())) groups.mu.Lock() delete(groups.freshErr, 2) @@ -537,8 +537,8 @@ func TestSchedulerFullRebuildPartialLifecycleFailureReturnsBeforeDBAndRetries(t require.NoError(t, svc.triggerFullRebuild("retry")) _, _, freshCalls = groups.stats() require.Equal(t, []int64{1, 2, 2, 3}, freshCalls) - require.Equal(t, 39, len(cache.retiredBuckets())) - require.Equal(t, 7, accounts.callCount()) + require.Equal(t, 57, len(cache.retiredBuckets())) + require.Equal(t, 10, accounts.callCount()) require.Empty(t, cache.tokens()) } @@ -561,14 +561,14 @@ func TestSchedulerFullRebuildActiveTombstoneLazyRecoveryDiscardsPartialCaptureTa capturesAtFirstDB = cache.captureAttemptCount() held, reopenCount := cache.leaseHeldAndTokenCount() require.False(t, held) - require.Equal(t, 12, reopenCount) - require.Equal(t, 19, capturesAtFirstDB) + require.Equal(t, 18, reopenCount) + require.Equal(t, 25, capturesAtFirstDB) } svc := newFullRebuildLifecycleService(cache, nil, accounts, groups, config.RunModeStandard) require.NoError(t, svc.rebuildFullSnapshot(context.Background(), "test")) require.Equal(t, capturesAtFirstDB, cache.captureAttemptCount()) - require.Equal(t, 15, accounts.callCount()) + require.Equal(t, 21, accounts.callCount()) for _, bucket := range canonical { attempts, published := cache.counts(bucket) require.Equal(t, 1, attempts, "discarded pre-recovery tokens must never publish: %s", bucket.String()) @@ -600,9 +600,9 @@ func TestSchedulerFullRebuildSimpleModePreservesRegistryWithoutLifecycleAuthorit require.Zero(t, activeCalls) require.Zero(t, fallbackCalls) require.Empty(t, freshCalls) - require.Equal(t, 15, cache.captureAttemptCount()) - require.Equal(t, 10, accounts.callCount()) - require.Equal(t, 10, accounts.groupCallCount(0)) + require.Equal(t, 21, cache.captureAttemptCount()) + require.Equal(t, 13, accounts.callCount()) + require.Equal(t, 13, accounts.groupCallCount(0)) require.Empty(t, cache.retiredBuckets()) require.Empty(t, cache.tokens()) for _, bucket := range registered { @@ -632,11 +632,11 @@ func TestSchedulerFullRebuildFreshReopenLockBusyRetriesWithoutBlockingOrdinaryTa require.Zero(t, cache.currentWatermark()) _, groupZeroPublished := cache.counts(schedulerCanonicalBuckets(0)[0]) require.Equal(t, 1, groupZeroPublished, "ordinary tasks must still run when one strict Reopen task is busy") - require.Equal(t, 14, accounts.callCount()) + require.Equal(t, 20, accounts.callCount()) svc.pollOutbox() require.Equal(t, int64(1), cache.currentWatermark()) - require.Equal(t, 28, accounts.callCount()) + require.Equal(t, 40, accounts.callCount()) _, busyBucketPublished := cache.counts(canonical[0]) require.Equal(t, 1, busyBucketPublished) activeCalls, fallbackCalls, freshCalls := groups.stats() @@ -654,7 +654,7 @@ func TestSchedulerFullRebuildOrdinaryLockBusyKeepsExistingSkipSemantics(t *testi svc := newFullRebuildLifecycleService(cache, nil, accounts, groups, config.RunModeStandard) require.NoError(t, svc.rebuildFullSnapshot(context.Background(), "test")) - require.Equal(t, 7, accounts.callCount()) + require.Equal(t, 10, accounts.callCount()) attempts, published := cache.counts(busyBucket) require.Zero(t, attempts) require.Zero(t, published) diff --git a/backend/internal/service/scheduler_snapshot_group_lifecycle_test.go b/backend/internal/service/scheduler_snapshot_group_lifecycle_test.go index ffc697eafc..f9e75ed158 100644 --- a/backend/internal/service/scheduler_snapshot_group_lifecycle_test.go +++ b/backend/internal/service/scheduler_snapshot_group_lifecycle_test.go @@ -325,8 +325,8 @@ func newGroupLifecycleTestService(cache SchedulerCache, accounts AccountReposito } func expectedGroupLifecycleBuckets(groupID int64) []SchedulerBucket { - platforms := []string{PlatformAnthropic, PlatformGemini, PlatformOpenAI, PlatformAntigravity, PlatformGrok} - buckets := make([]SchedulerBucket, 0, 12) + platforms := schedulerSnapshotPlatforms() + buckets := make([]SchedulerBucket, 0, 18) for _, platform := range platforms { buckets = append(buckets, SchedulerBucket{GroupID: groupID, Platform: platform, Mode: SchedulerModeSingle}, @@ -441,7 +441,7 @@ func TestSchedulerGroupLifecycleActiveReopensAndRebuildsAllCurrentBuckets(t *tes accounts.beforeLoad = func() { held, tokenCount := cache.leaseHeldAndTokenCount() require.False(t, held, "the group lifecycle lease must be released before the first account query") - require.Equal(t, 12, tokenCount, "all reopen tokens must be prepared before the first account query") + require.Equal(t, 18, tokenCount, "all reopen tokens must be prepared before the first account query") } svc := newGroupLifecycleTestService(cache, accounts, groups, config.RunModeStandard) seen := make(map[batchSeenKey]struct{}) @@ -453,8 +453,8 @@ func TestSchedulerGroupLifecycleActiveReopensAndRebuildsAllCurrentBuckets(t *tes registered, err := cache.retirementRaceCache.ListBuckets(context.Background()) require.NoError(t, err) require.Contains(t, bucketStrings(registered), historical.String()) - require.Len(t, cache.tokens(), 12) - require.Equal(t, 7, accounts.callCount()) + require.Len(t, cache.tokens(), 18) + require.Equal(t, 10, accounts.callCount()) require.Equal(t, 1, accounts.platformCallCount(PlatformOpenAI)) for _, bucket := range current { _, published := cache.counts(bucket) @@ -472,16 +472,16 @@ func TestSchedulerGroupLifecycleActiveReopensAndRebuildsAllCurrentBuckets(t *tes require.True(t, cache.releaseDeadline) require.NoError(t, cache.releaseCtxErr) _, reopenHeld := cache.lifecycleMutationLeaseStates() - require.Len(t, reopenHeld, 12) + require.Len(t, reopenHeld, 18) for _, held := range reopenHeld { require.True(t, held) } lockTTLs, unlockCalls := cache.lockStats() - require.Len(t, lockTTLs, 12) + require.Len(t, lockTTLs, 18) for _, ttl := range lockTTLs { require.Equal(t, 30*time.Second, ttl) } - require.Equal(t, 12, unlockCalls) + require.Equal(t, 18, unlockCalls) requireLifecycleSeen(t, seen, groupID) } @@ -497,8 +497,8 @@ func TestSchedulerGroupLifecycleInactiveThenActiveAuthoritativelyReopens(t *test groups.set(&Group{ID: groupID, Status: StatusActive, Hydrated: true}, nil) require.NoError(t, svc.handleGroupEvent(context.Background(), ptrInt64(groupID), make(map[batchSeenKey]struct{}))) - require.Len(t, cache.tokens(), 12) - require.Equal(t, 7, accounts.callCount()) + require.Len(t, cache.tokens(), 18) + require.Equal(t, 10, accounts.callCount()) for _, bucket := range expectedGroupLifecycleBuckets(groupID) { _, published := cache.counts(bucket) require.Equal(t, 1, published, bucket.String()) @@ -546,15 +546,15 @@ func TestSchedulerGroupLifecycleEpochPreventsABA(t *testing.T) { groups.set(&Group{ID: groupID, Status: StatusActive, Hydrated: true}, nil) require.NoError(t, svc.handleGroupEvent(context.Background(), ptrInt64(groupID), make(map[batchSeenKey]struct{}))) firstActiveTokens := cache.tokens() - require.Len(t, firstActiveTokens, 12) + require.Len(t, firstActiveTokens, 18) groups.set(&Group{ID: groupID, Status: StatusDisabled, Hydrated: true}, nil) require.NoError(t, svc.handleGroupEvent(context.Background(), ptrInt64(groupID), make(map[batchSeenKey]struct{}))) groups.set(&Group{ID: groupID, Status: StatusActive, Hydrated: true}, nil) require.NoError(t, svc.handleGroupEvent(context.Background(), ptrInt64(groupID), make(map[batchSeenKey]struct{}))) allTokens := cache.tokens() - require.Len(t, allTokens, 24) - require.Greater(t, allTokens[12].Epoch, firstActiveTokens[0].Epoch) + require.Len(t, allTokens, 36) + require.Greater(t, allTokens[18].Epoch, firstActiveTokens[0].Epoch) require.ErrorIs(t, cache.SetSnapshot(context.Background(), firstActiveTokens[0].Bucket, firstActiveTokens[0], nil), ErrSchedulerBucketWriteFenced) } @@ -571,11 +571,11 @@ func TestSchedulerGroupLifecycleSeenIsIndependentAndDeduplicatesGroupEvents(t *t require.NoError(t, svc.handleGroupEvent(context.Background(), ptrInt64(groupID), seen)) require.Equal(t, 1, groups.callCount()) - require.Equal(t, 7, accounts.callCount()) + require.Equal(t, 10, accounts.callCount()) requireLifecycleSeen(t, seen, groupID) require.NoError(t, svc.handleGroupEvent(context.Background(), ptrInt64(groupID), seen)) require.Equal(t, 1, groups.callCount()) - require.Equal(t, 7, accounts.callCount()) + require.Equal(t, 10, accounts.callCount()) } func TestSchedulerGroupLifecycleFailuresDoNotMarkSeen(t *testing.T) { diff --git a/backend/internal/service/scheduler_snapshot_retirement_test.go b/backend/internal/service/scheduler_snapshot_retirement_test.go index 3d24179758..6fa2b7ddc5 100644 --- a/backend/internal/service/scheduler_snapshot_retirement_test.go +++ b/backend/internal/service/scheduler_snapshot_retirement_test.go @@ -176,7 +176,7 @@ func TestSchedulerFullRebuildCapturesAllRegistryTokensBeforeDBLoad(t *testing.T) } captures, reopens := cache.captureAndReopenCounts() - require.Equal(t, 24, captures, "group0 and active-group canonical tokens must be captured before the first DB load") + require.Equal(t, 36, captures, "group0 and active-group canonical tokens must be captured before the first DB load") require.Zero(t, reopens) require.NoError(t, cache.RetireBucket(context.Background(), queued)) _, err := cache.ReopenBucket(context.Background(), queued) diff --git a/backend/internal/service/scheduler_snapshot_service.go b/backend/internal/service/scheduler_snapshot_service.go index 9655f54f25..e38d0291ba 100644 --- a/backend/internal/service/scheduler_snapshot_service.go +++ b/backend/internal/service/scheduler_snapshot_service.go @@ -587,7 +587,7 @@ func (s *SchedulerSnapshotService) handleBulkAccountEvent(ctx context.Context, p rebuildGroupIDs = append(rebuildGroupIDs, gid) } - // 缺失账户无法确定原平台,保留五平台重建以避免遗留旧快照。 + // 缺失账户无法确定原平台,保留全平台重建以避免遗留旧快照。 if !allAccountsFound { return s.rebuildByGroupIDs(ctx, rebuildGroupIDs, "account_bulk_change", seen) } @@ -609,7 +609,7 @@ func (s *SchedulerSnapshotService) handleBulkAccountEvent(ctx context.Context, p } accountGroupIDs := s.normalizeGroupIDs(account.GroupIDs) switch account.Platform { - case PlatformAnthropic, PlatformGemini, PlatformOpenAI, PlatformGrok: + case PlatformAnthropic, PlatformGemini, PlatformOpenAI, PlatformGrok, PlatformKimi, PlatformZhipu, PlatformDeepseek: addPlatformGroups(account.Platform, accountGroupIDs) case PlatformAntigravity: // 批量更新可能刚关闭 mixed_scheduling,仍需清理两个兼容平台的旧快照。 @@ -824,8 +824,8 @@ func (s *SchedulerSnapshotService) rebuildByAccount(ctx context.Context, account return s.rebuildBuckets(ctx, buckets, reason) } -func schedulerSnapshotPlatforms() [5]string { - return [5]string{PlatformAnthropic, PlatformGemini, PlatformOpenAI, PlatformAntigravity, PlatformGrok} +func schedulerSnapshotPlatforms() [8]string { + return [8]string{PlatformAnthropic, PlatformGemini, PlatformOpenAI, PlatformAntigravity, PlatformGrok, PlatformKimi, PlatformZhipu, PlatformDeepseek} } // 生命周期辅助函数有意排除 group0;full rebuild 构造 group0 canonical 集时必须显式调用 canonical helper。 @@ -837,7 +837,7 @@ func schedulerBucketsForGroup(groupID int64) []SchedulerBucket { } func schedulerCanonicalBuckets(groupID int64) []SchedulerBucket { - buckets := make([]SchedulerBucket, 0, 12) + buckets := make([]SchedulerBucket, 0, 18) for _, platform := range schedulerSnapshotPlatforms() { buckets = append(buckets, SchedulerBucket{GroupID: groupID, Platform: platform, Mode: SchedulerModeSingle}, @@ -855,7 +855,7 @@ func (s *SchedulerSnapshotService) rebuildByGroupIDs(ctx context.Context, groupI if len(groupIDs) == 0 { return nil } - buckets := make([]SchedulerBucket, 0, len(groupIDs)*12) + buckets := make([]SchedulerBucket, 0, len(groupIDs)*18) for _, platform := range schedulerSnapshotPlatforms() { buckets = append(buckets, s.bucketsForPlatform(platform, groupIDs, seen)...) } diff --git a/backend/migrations/227_composite_routes_add_cn_providers.sql b/backend/migrations/227_composite_routes_add_cn_providers.sql new file mode 100644 index 0000000000..ec4b228241 --- /dev/null +++ b/backend/migrations/227_composite_routes_add_cn_providers.sql @@ -0,0 +1,8 @@ +-- Allow Composite model routes to target the three concrete CN providers. +ALTER TABLE composite_model_routes + DROP CONSTRAINT IF EXISTS composite_model_routes_target_platform_check; + +ALTER TABLE composite_model_routes + ADD CONSTRAINT composite_model_routes_target_platform_check + CHECK (target_platform IN ('anthropic', 'openai', 'gemini', 'antigravity', 'grok', + 'kimi', 'zhipu', 'deepseek')); diff --git a/backend/migrations/composite_routes_cn_providers_migration_test.go b/backend/migrations/composite_routes_cn_providers_migration_test.go new file mode 100644 index 0000000000..7b6bba7471 --- /dev/null +++ b/backend/migrations/composite_routes_cn_providers_migration_test.go @@ -0,0 +1,18 @@ +package migrations + +import ( + "strings" + "testing" + + "github.com/stretchr/testify/require" +) + +func TestCompositeRoutesCNProvidersMigration(t *testing.T) { + content, err := FS.ReadFile("227_composite_routes_add_cn_providers.sql") + require.NoError(t, err) + + sql := strings.Join(strings.Fields(string(content)), " ") + require.Contains(t, sql, "DROP CONSTRAINT IF EXISTS composite_model_routes_target_platform_check") + require.Contains(t, sql, + "CHECK (target_platform IN ('anthropic', 'openai', 'gemini', 'antigravity', 'grok', 'kimi', 'zhipu', 'deepseek'))") +} diff --git a/frontend/src/views/admin/ChannelsView.vue b/frontend/src/views/admin/ChannelsView.vue index e3a7ca8262..4da09f900c 100644 --- a/frontend/src/views/admin/ChannelsView.vue +++ b/frontend/src/views/admin/ChannelsView.vue @@ -763,9 +763,8 @@ let abortController: AbortController | null = null // ── Platform config ── const platformOrder: GroupPlatform[] = ['anthropic', 'openai', 'gemini', 'antigravity', 'grok', 'kimi', 'zhipu', 'deepseek'] -// composite 分组仅覆盖主平台(与后端 isConcreteRequestPlatform / composite-routes target_platform 一致), -// 不含国产供应商平台。 -const compositePlatforms: GroupPlatform[] = ['anthropic', 'openai', 'gemini', 'antigravity', 'grok'] +// Composite pricing/mapping may target every concrete schedulable provider. +const compositePlatforms: GroupPlatform[] = ['anthropic', 'openai', 'gemini', 'antigravity', 'grok', 'kimi', 'zhipu', 'deepseek'] // ── Helpers ── function formatDate(value: string): string { diff --git a/frontend/src/views/admin/GroupsView.vue b/frontend/src/views/admin/GroupsView.vue index b0144f4686..3fff44b131 100644 --- a/frontend/src/views/admin/GroupsView.vue +++ b/frontend/src/views/admin/GroupsView.vue @@ -4771,6 +4771,9 @@ const compositeRoutePlatformOptions = computed(() => [ { value: "gemini", label: "Gemini" }, { value: "antigravity", label: "Antigravity" }, { value: "grok", label: "Grok" }, + { value: "kimi", label: "Kimi" }, + { value: "zhipu", label: "Zhipu GLM" }, + { value: "deepseek", label: "DeepSeek" }, ]); const compositeRouteEndpointOptions = computed(() => [ diff --git a/frontend/src/views/admin/__tests__/GroupsView.compositePlatforms.spec.ts b/frontend/src/views/admin/__tests__/GroupsView.compositePlatforms.spec.ts new file mode 100644 index 0000000000..d6b5abcfd7 --- /dev/null +++ b/frontend/src/views/admin/__tests__/GroupsView.compositePlatforms.spec.ts @@ -0,0 +1,17 @@ +import { readFileSync } from 'node:fs' +import { resolve } from 'node:path' +import { describe, expect, it } from 'vitest' + +describe('GroupsView Composite route options', () => { + it('offers Kimi, Zhipu GLM, and DeepSeek as route targets', () => { + const source = readFileSync(resolve('src/views/admin/GroupsView.vue'), 'utf8') + const options = source.slice( + source.indexOf('const compositeRoutePlatformOptions'), + source.indexOf('const compositeRouteEndpointOptions') + ) + + expect(options).toContain('{ value: "kimi", label: "Kimi" }') + expect(options).toContain('{ value: "zhipu", label: "Zhipu GLM" }') + expect(options).toContain('{ value: "deepseek", label: "DeepSeek" }') + }) +}) diff --git a/frontend/src/views/admin/__tests__/channelPlatformOptions.spec.ts b/frontend/src/views/admin/__tests__/channelPlatformOptions.spec.ts new file mode 100644 index 0000000000..bd078c8175 --- /dev/null +++ b/frontend/src/views/admin/__tests__/channelPlatformOptions.spec.ts @@ -0,0 +1,14 @@ +import { readFileSync } from 'node:fs' +import { resolve } from 'node:path' +import { describe, expect, it } from 'vitest' + +describe('Composite channel platform options', () => { + it('includes the CN concrete providers for pricing and model mapping', () => { + const source = readFileSync(resolve('src/views/admin/ChannelsView.vue'), 'utf8') + const declaration = source.match(/const compositePlatforms:[^=]+=[^\n]+/)?.[0] + + expect(declaration).toContain("'kimi'") + expect(declaration).toContain("'zhipu'") + expect(declaration).toContain("'deepseek'") + }) +})