mirror of
https://github.com/Wei-Shaw/sub2api.git
synced 2026-10-07 15:08:02 +08:00
Merge pull request #5817 from wucm667/fix/issue-5796-composite-new-platforms
fix(composite): support Kimi, GLM, and DeepSeek
This commit is contained in:
@@ -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)
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -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"))
|
||||
})
|
||||
}
|
||||
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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'))")
|
||||
}
|
||||
@@ -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 {
|
||||
|
||||
@@ -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'")
|
||||
})
|
||||
})
|
||||
Reference in New Issue
Block a user