Merge pull request #5817 from wucm667/fix/issue-5796-composite-new-platforms

fix(composite): support Kimi, GLM, and DeepSeek
This commit is contained in:
Wesley Liddick
2026-08-19 15:36:38 +08:00
committed by GitHub
26 changed files with 398 additions and 100 deletions
@@ -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"`
@@ -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)
}
}
@@ -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)
}
}
+5 -3
View File
@@ -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
@@ -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)
@@ -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)
}
}
@@ -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)
@@ -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
}
@@ -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"))
})
}
+1 -1
View File
@@ -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
@@ -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)
}
}
+1 -1
View File
@@ -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}
}
@@ -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 {
+15 -1
View File
@@ -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
@@ -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))
}
}
@@ -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})
@@ -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)
@@ -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) {
@@ -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)
@@ -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)...)
}
@@ -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'));
@@ -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'))")
}
+2 -3
View File
@@ -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 {
+3
View File
@@ -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(() => [
@@ -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" }')
})
})
@@ -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'")
})
})