diff --git a/README.md b/README.md index f47968421d..0e87f75997 100644 --- a/README.md +++ b/README.md @@ -204,6 +204,7 @@ Sub2API is an AI API gateway platform designed to distribute and manage API quot - **Rate Limiting** - Configurable request and token rate limits - **Built-in Payment System** - Supports EasyPay, Alipay, WeChat Pay, and Stripe for user self-service top-up, no separate payment service needed ([Configuration Guide](docs/PAYMENT.md)) - **Admin Dashboard** - Web interface for monitoring and management +- **Composite Groups** - Admin routing layer that resolves requested models to concrete providers for multi-provider groups ([Operator Guide](docs/COMPOSITE_GROUPS.md)) - **External System Integration** - Embed external systems (e.g. ticketing) via iframe to extend the admin dashboard ## Ecosystem diff --git a/backend/internal/handler/composite_platform.go b/backend/internal/handler/composite_platform.go index b2f72cabab..bb001a8b22 100644 --- a/backend/internal/handler/composite_platform.go +++ b/backend/internal/handler/composite_platform.go @@ -35,6 +35,15 @@ func compositeTargetPlatformAllowed(c *gin.Context, apiKey *service.APIKey, mode return false } +func compositeTargetPlatformResolved(c *gin.Context, apiKey *service.APIKey, model string) bool { + if c == nil || c.Request == nil || apiKey == nil || apiKey.Group == nil || apiKey.Group.Platform != service.PlatformComposite { + return true + } + ensureCompositeTargetPlatform(c, apiKey, model) + _, ok := service.ResolvedTargetPlatformFromContext(c.Request.Context()) + return ok +} + func effectiveAPIKeyPlatform(c *gin.Context, apiKey *service.APIKey) string { if c != nil && c.Request != nil { if platform, ok := service.ResolvedTargetPlatformFromContext(c.Request.Context()); ok { diff --git a/backend/internal/handler/composite_platform_test.go b/backend/internal/handler/composite_platform_test.go index 46a2595ebb..dc80ca3ff6 100644 --- a/backend/internal/handler/composite_platform_test.go +++ b/backend/internal/handler/composite_platform_test.go @@ -40,3 +40,23 @@ func TestCompositeTargetPlatformAllowedRejectsWrongOrUnknownModel(t *testing.T) }) } } + +func TestCompositeTargetPlatformResolvedRejectsUnknownModel(t *testing.T) { + gin.SetMode(gin.TestMode) + c, _ := gin.CreateTestContext(httptest.NewRecorder()) + c.Request = httptest.NewRequest("POST", "/v1/messages", nil) + apiKey := &service.APIKey{Group: &service.Group{Platform: service.PlatformComposite}} + + require.False(t, compositeTargetPlatformResolved(c, apiKey, "llama-4-maverick")) + _, ok := service.ResolvedTargetPlatformFromContext(c.Request.Context()) + require.False(t, ok) +} + +func TestCompositeTargetPlatformResolvedAllowsConcreteGroupWithoutResolution(t *testing.T) { + gin.SetMode(gin.TestMode) + c, _ := gin.CreateTestContext(httptest.NewRecorder()) + c.Request = httptest.NewRequest("POST", "/v1/messages", nil) + apiKey := &service.APIKey{Group: &service.Group{Platform: service.PlatformAnthropic}} + + require.True(t, compositeTargetPlatformResolved(c, apiKey, "llama-4-maverick")) +} diff --git a/backend/internal/handler/gateway_handler.go b/backend/internal/handler/gateway_handler.go index 5ed98e6ec8..739b9e8e44 100644 --- a/backend/internal/handler/gateway_handler.go +++ b/backend/internal/handler/gateway_handler.go @@ -202,6 +202,10 @@ func (h *GatewayHandler) Messages(c *gin.Context) { h.errorResponse(c, http.StatusBadRequest, "invalid_request_error", "model is required") return } + if !compositeTargetPlatformResolved(c, apiKey, reqModel) { + h.errorResponse(c, http.StatusBadRequest, "invalid_request_error", "Model is not supported by composite groups") + return + } if decision := h.checkSecurityAudit(c, reqLog, apiKey, subject, service.ContentModerationProtocolAnthropicMessages, reqModel, body); decision != nil && !decision.AllowNextStage { h.anthropicSecurityAuditError(c, decision) @@ -1021,6 +1025,11 @@ func (h *GatewayHandler) Models(c *gin.Context) { if platform == service.PlatformComposite { availableModels := h.compositeAvailableModels(c.Request.Context(), groupID) + if apiKey != nil && apiKey.Group != nil && apiKey.Group.CustomModelsListEnabled() { + availableModels = filterModelsByCustomList(availableModels, defaultModelIDsForPlatform(service.PlatformComposite), apiKey.Group.ModelsListConfig.Models) + writeCustomModelsList(c, service.PlatformComposite, availableModels) + return + } if len(availableModels) > 0 { writeModelsList(c, service.PlatformComposite, availableModels) return @@ -1947,6 +1956,10 @@ func (h *GatewayHandler) CountTokens(c *gin.Context) { h.errorResponse(c, http.StatusBadRequest, "invalid_request_error", "model is required") return } + if !compositeTargetPlatformResolved(c, apiKey, parsedReq.Model) { + h.errorResponse(c, http.StatusBadRequest, "invalid_request_error", "Model is not supported by composite groups") + return + } setOpsRequestContext(c, parsedReq.Model, parsedReq.Stream) setOpsEndpointContext(c, "", int16(service.RequestTypeFromLegacy(parsedReq.Stream, false))) diff --git a/backend/internal/handler/gateway_handler_chat_completions.go b/backend/internal/handler/gateway_handler_chat_completions.go index 51ddf8bcac..0f779441d0 100644 --- a/backend/internal/handler/gateway_handler_chat_completions.go +++ b/backend/internal/handler/gateway_handler_chat_completions.go @@ -76,6 +76,10 @@ func (h *GatewayHandler) ChatCompletions(c *gin.Context) { } reqModel := modelResult.String() ensureCompositeTargetPlatform(c, apiKey, reqModel) + if !compositeTargetPlatformResolved(c, apiKey, reqModel) { + h.chatCompletionsErrorResponse(c, http.StatusBadRequest, "invalid_request_error", "Model is not supported by composite groups") + return + } reqStream, ok := parseOpenAICompatibleStream(body) if !ok { h.chatCompletionsErrorResponse(c, http.StatusBadRequest, "invalid_request_error", invalidStreamFieldTypeMessage) diff --git a/backend/internal/handler/gateway_handler_responses.go b/backend/internal/handler/gateway_handler_responses.go index 7a96a611c7..475ad9e88f 100644 --- a/backend/internal/handler/gateway_handler_responses.go +++ b/backend/internal/handler/gateway_handler_responses.go @@ -76,6 +76,10 @@ func (h *GatewayHandler) Responses(c *gin.Context) { } reqModel := modelResult.String() ensureCompositeTargetPlatform(c, apiKey, reqModel) + if !compositeTargetPlatformResolved(c, apiKey, reqModel) { + h.responsesErrorResponse(c, http.StatusBadRequest, "invalid_request_error", "Model is not supported by composite groups") + return + } reqStream, ok := parseOpenAICompatibleStream(body) if !ok { h.responsesErrorResponse(c, http.StatusBadRequest, "invalid_request_error", invalidStreamFieldTypeMessage) diff --git a/backend/internal/handler/gateway_models_test.go b/backend/internal/handler/gateway_models_test.go index 30a8b5463f..72fe837f70 100644 --- a/backend/internal/handler/gateway_models_test.go +++ b/backend/internal/handler/gateway_models_test.go @@ -277,6 +277,61 @@ func TestGatewayModels_CustomModelsListFiltersAndOrdersMappedModels(t *testing.T require.Equal(t, []string{"gpt-5.5", "gpt-5.4"}, modelIDsForTest(got.Data)) } +func TestGatewayModels_CompositeCustomModelsListFiltersAcrossConcretePlatforms(t *testing.T) { + gin.SetMode(gin.TestMode) + + groupID := int64(33) + h := newGatewayModelsHandlerForTest( + &gatewayModelsAccountRepoStub{ + byGroup: map[int64][]service.Account{ + groupID: { + { + ID: 1, + Platform: service.PlatformOpenAI, + Credentials: map[string]any{ + "model_mapping": map[string]any{ + "gpt-5.4": "gpt-5.4", + "gpt-5.5": "gpt-5.5", + }, + }, + }, + { + ID: 2, + Platform: service.PlatformGemini, + Credentials: map[string]any{ + "model_mapping": map[string]any{ + "gemini-2.5-flash": "gemini-2.5-flash", + }, + }, + }, + }, + }, + }, + ) + + 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, + ModelsListConfig: service.GroupModelsListConfig{ + Enabled: true, + Models: []string{"gemini-2.5-flash", "missing-model", "gpt-5.5"}, + }, + }, + }) + + h.Models(c) + + require.Equal(t, http.StatusOK, rec.Code) + + var got gatewayModelsResponseForTest + require.NoError(t, json.Unmarshal(rec.Body.Bytes(), &got)) + require.Equal(t, []string{"gemini-2.5-flash", "gpt-5.5"}, modelIDsForTest(got.Data)) +} + func TestGatewayModels_CustomModelsListKeepsConcreteModelAllowedByWildcardMapping(t *testing.T) { gin.SetMode(gin.TestMode) diff --git a/backend/internal/handler/openai_gateway_handler.go b/backend/internal/handler/openai_gateway_handler.go index 6aa57a3f45..8c24c85172 100644 --- a/backend/internal/handler/openai_gateway_handler.go +++ b/backend/internal/handler/openai_gateway_handler.go @@ -2798,7 +2798,11 @@ func (h *OpenAIGatewayHandler) enqueueCyberSessionBlockedOpsEntry(c *gin.Context meta.Stream = b } } - meta.Platform = resolveOpsPlatform(apiKey, guessPlatformFromPath(meta.RequestPath)) + requestCtx := context.Background() + if c.Request != nil { + requestCtx = c.Request.Context() + } + meta.Platform = resolveOpsPlatform(requestCtx, apiKey, guessPlatformFromPath(meta.RequestPath)) if c.Request != nil { meta.ClientRequestID, _ = c.Request.Context().Value(ctxkey.ClientRequestID).(string) meta.UserAgent = c.GetHeader("User-Agent") @@ -2864,7 +2868,11 @@ func (h *OpenAIGatewayHandler) recordCyberPolicyIfMarked(c *gin.Context, apiKey if c.Request != nil && c.Request.URL != nil { requestPath = c.Request.URL.Path } - platform := resolveOpsPlatform(apiKey, guessPlatformFromPath(requestPath)) + requestCtx := context.Background() + if c.Request != nil { + requestCtx = c.Request.Context() + } + platform := resolveOpsPlatform(requestCtx, apiKey, guessPlatformFromPath(requestPath)) var clientRequestID, userAgent, clientIPStr string if c.Request != nil { clientRequestID, _ = c.Request.Context().Value(ctxkey.ClientRequestID).(string) diff --git a/backend/internal/handler/ops_error_logger.go b/backend/internal/handler/ops_error_logger.go index 6f09e64084..5eb595f2ba 100644 --- a/backend/internal/handler/ops_error_logger.go +++ b/backend/internal/handler/ops_error_logger.go @@ -784,7 +784,7 @@ func OpsErrorLoggerMiddleware(ops *service.OpsService) gin.HandlerFunc { } fallbackPlatform := guessPlatformFromPath(c.Request.URL.Path) - platform := resolveOpsPlatform(apiKey, fallbackPlatform) + platform := resolveOpsPlatform(c.Request.Context(), apiKey, fallbackPlatform) requestID := c.Writer.Header().Get("X-Request-Id") if requestID == "" { @@ -1005,7 +1005,7 @@ func OpsErrorLoggerMiddleware(ops *service.OpsService) gin.HandlerFunc { } fallbackPlatform := guessPlatformFromPath(c.Request.URL.Path) - platform := resolveOpsPlatform(apiKey, fallbackPlatform) + platform := resolveOpsPlatform(c.Request.Context(), apiKey, fallbackPlatform) requestID := c.Writer.Header().Get("X-Request-Id") if requestID == "" { @@ -1431,7 +1431,10 @@ func getOpsAPIKey(c *gin.Context) *service.APIKey { return nil } -func resolveOpsPlatform(apiKey *service.APIKey, fallback string) string { +func resolveOpsPlatform(ctx context.Context, apiKey *service.APIKey, fallback string) string { + if platform, ok := service.ResolvedTargetPlatformFromContext(ctx); ok { + return platform + } if apiKey != nil && apiKey.Group != nil && apiKey.Group.Platform != "" { return apiKey.Group.Platform } diff --git a/backend/internal/handler/ops_platform_test.go b/backend/internal/handler/ops_platform_test.go new file mode 100644 index 0000000000..9053b3bcab --- /dev/null +++ b/backend/internal/handler/ops_platform_test.go @@ -0,0 +1,16 @@ +package handler + +import ( + "context" + "testing" + + "github.com/Wei-Shaw/sub2api/internal/service" + "github.com/stretchr/testify/require" +) + +func TestResolveOpsPlatformPrefersResolvedCompositeTarget(t *testing.T) { + apiKey := &service.APIKey{Group: &service.Group{Platform: service.PlatformComposite}} + ctx := service.WithResolvedTargetPlatform(context.Background(), service.PlatformOpenAI) + + require.Equal(t, service.PlatformOpenAI, resolveOpsPlatform(ctx, apiKey, service.PlatformAnthropic)) +} diff --git a/backend/internal/repository/usage_log_effective_platform_test.go b/backend/internal/repository/usage_log_effective_platform_test.go new file mode 100644 index 0000000000..a800550844 --- /dev/null +++ b/backend/internal/repository/usage_log_effective_platform_test.go @@ -0,0 +1,16 @@ +package repository + +import ( + "strings" + "testing" + + "github.com/stretchr/testify/require" +) + +func TestUsageLogEffectivePlatformExprUsesAccountPlatformForCompositeGroups(t *testing.T) { + expr := strings.ToLower(usageLogEffectivePlatformExpr) + + require.Contains(t, expr, "g.platform = 'composite'") + require.Contains(t, expr, "then a.platform") + require.Contains(t, expr, "coalesce") +} diff --git a/backend/internal/repository/usage_log_repo.go b/backend/internal/repository/usage_log_repo.go index 4359edc2ca..94b605693d 100644 --- a/backend/internal/repository/usage_log_repo.go +++ b/backend/internal/repository/usage_log_repo.go @@ -31,8 +31,10 @@ const usageLogSuccessFilterUL = "ul.actual_cost > 0" // usageLogEffectivePlatformExpr 用于按"有效平台"维度聚合 usage_logs: // 优先取请求实际走的分组 platform,若分组未设置 platform 再 fallback 到 account.platform。 +// Composite groups are a routing layer, so platform analytics must use the +// resolved concrete account platform instead of grouping spend under "composite". // 配套要求查询里 LEFT JOIN groups g ON g.id = ul.group_id 与 LEFT JOIN accounts a ON a.id = ul.account_id。 -const usageLogEffectivePlatformExpr = "COALESCE(NULLIF(g.platform,''), a.platform)" +const usageLogEffectivePlatformExpr = "CASE WHEN g.platform = 'composite' THEN a.platform ELSE COALESCE(NULLIF(g.platform,''), a.platform) END" // dateFormatWhitelist 将 granularity 参数映射为 PostgreSQL TO_CHAR 格式字符串,防止外部输入直接拼入 SQL var dateFormatWhitelist = map[string]string{ diff --git a/backend/internal/service/admin_group.go b/backend/internal/service/admin_group.go index 80eee93617..d2e54e31b6 100644 --- a/backend/internal/service/admin_group.go +++ b/backend/internal/service/admin_group.go @@ -78,7 +78,11 @@ func (s *adminServiceImpl) GetGroupModelsListCandidates(ctx context.Context, id seen[model] = struct{}{} } for _, acc := range accounts { - if acc.Platform != platform { + if platform == PlatformComposite { + if !isConcreteRequestPlatform(acc.Platform) { + continue + } + } else if acc.Platform != platform { continue } for model := range acc.GetModelMapping() { @@ -116,7 +120,7 @@ func defaultModelsListCandidateIDs(platform string) []string { case PlatformGrok: return xai.DefaultModelIDs() case PlatformComposite: - return nil + return compositeDefaultModelsListCandidateIDs() default: ids := make([]string, 0, len(claude.DefaultModels)) for _, model := range claude.DefaultModels { @@ -132,6 +136,21 @@ func defaultAllowImageGenerationForPlatform(platform string) bool { return platform == PlatformGrok } +func compositeDefaultModelsListCandidateIDs() []string { + seen := make(map[string]struct{}) + ids := make([]string, 0) + for _, platform := range []string{PlatformAnthropic, PlatformGemini, PlatformOpenAI, PlatformAntigravity, PlatformGrok} { + for _, id := range defaultModelsListCandidateIDs(platform) { + if _, ok := seen[id]; ok { + continue + } + seen[id] = struct{}{} + ids = append(ids, id) + } + } + return ids +} + func canCopyAccountsFromGroupPlatform(targetPlatform, sourcePlatform string) bool { if targetPlatform == PlatformComposite { return sourcePlatform == PlatformComposite || isConcreteRequestPlatform(sourcePlatform) diff --git a/backend/internal/service/admin_service_bulk_update_test.go b/backend/internal/service/admin_service_bulk_update_test.go index 2f44b1741d..63bf26fbd1 100644 --- a/backend/internal/service/admin_service_bulk_update_test.go +++ b/backend/internal/service/admin_service_bulk_update_test.go @@ -14,25 +14,31 @@ import ( type accountRepoStubForBulkUpdate struct { accountRepoStub - bulkUpdateErr error - bulkUpdateIDs []int64 - bindGroupErrByID map[int64]error - bindGroupsCalls []int64 - getByIDsAccounts []*Account - getByIDsErr error - getByIDsCalled bool - getByIDsIDs []int64 - getByIDAccounts map[int64]*Account - getByIDErrByID map[int64]error - getByIDCalled []int64 - listByGroupData map[int64][]Account - listByGroupErr map[int64]error - listData []Account - listResult *pagination.PaginationResult - listErr error - listCalled bool - lastListParams pagination.PaginationParams - lastListFilters struct { + bulkUpdateErr error + bulkUpdateIDs []int64 + bindGroupErrByID map[int64]error + bindGroupsCalls []int64 + bindGroupsByAccount map[int64][]int64 + createAccount *Account + createID int64 + createErr error + updatedAccounts []*Account + updateErr error + getByIDsAccounts []*Account + getByIDsErr error + getByIDsCalled bool + getByIDsIDs []int64 + getByIDAccounts map[int64]*Account + getByIDErrByID map[int64]error + getByIDCalled []int64 + listByGroupData map[int64][]Account + listByGroupErr map[int64]error + listData []Account + listResult *pagination.PaginationResult + listErr error + listCalled bool + lastListParams pagination.PaginationParams + lastListFilters struct { platform string accountType string status string @@ -50,8 +56,25 @@ func (s *accountRepoStubForBulkUpdate) BulkUpdate(_ context.Context, ids []int64 return int64(len(ids)), nil } -func (s *accountRepoStubForBulkUpdate) BindGroups(_ context.Context, accountID int64, _ []int64) error { +func (s *accountRepoStubForBulkUpdate) Create(_ context.Context, account *Account) error { + s.createAccount = account + if s.createID > 0 { + account.ID = s.createID + } + return s.createErr +} + +func (s *accountRepoStubForBulkUpdate) Update(_ context.Context, account *Account) error { + s.updatedAccounts = append(s.updatedAccounts, account) + return s.updateErr +} + +func (s *accountRepoStubForBulkUpdate) BindGroups(_ context.Context, accountID int64, groupIDs []int64) error { s.bindGroupsCalls = append(s.bindGroupsCalls, accountID) + if s.bindGroupsByAccount == nil { + s.bindGroupsByAccount = make(map[int64][]int64) + } + s.bindGroupsByAccount[accountID] = append([]int64{}, groupIDs...) if err, ok := s.bindGroupErrByID[accountID]; ok { return err } diff --git a/backend/internal/service/admin_service_composite_group_test.go b/backend/internal/service/admin_service_composite_group_test.go new file mode 100644 index 0000000000..73095abc5f --- /dev/null +++ b/backend/internal/service/admin_service_composite_group_test.go @@ -0,0 +1,181 @@ +//go:build unit + +package service + +import ( + "context" + "testing" + + "github.com/stretchr/testify/require" +) + +type accountRepoStubForCompositeModelsList struct { + accountRepoStub + accounts []Account +} + +func (s *accountRepoStubForCompositeModelsList) ListSchedulableByGroupID(_ context.Context, _ int64) ([]Account, error) { + return s.accounts, nil +} + +func TestAdminService_CreateCompositeGroupCopiesAccountsFromConcreteGroups(t *testing.T) { + var copiedFrom []int64 + var boundGroupID int64 + var boundAccountIDs []int64 + groupRepo := &groupRepoStubForAdmin{ + createID: 99, + getByIDByID: map[int64]*Group{ + 10: {ID: 10, Platform: PlatformOpenAI}, + 20: {ID: 20, Platform: PlatformGemini}, + }, + getAccountIDsByGroupIDsFn: func(groupIDs []int64) ([]int64, error) { + copiedFrom = append([]int64{}, groupIDs...) + return []int64{101, 202}, nil + }, + bindAccountsToGroupFn: func(groupID int64, accountIDs []int64) error { + boundGroupID = groupID + boundAccountIDs = append([]int64{}, accountIDs...) + return nil + }, + } + svc := &adminServiceImpl{groupRepo: groupRepo} + + group, err := svc.CreateGroup(context.Background(), &CreateGroupInput{ + Name: "Composite", + Platform: PlatformComposite, + RateMultiplier: 1, + CopyAccountsFromGroupIDs: []int64{10, 20, 10}, + }) + + require.NoError(t, err) + require.Equal(t, PlatformComposite, groupRepo.created.Platform) + require.Equal(t, int64(99), group.ID) + require.Equal(t, int64(2), group.AccountCount) + require.ElementsMatch(t, []int64{10, 20}, copiedFrom) + require.Equal(t, int64(99), boundGroupID) + require.ElementsMatch(t, []int64{101, 202}, boundAccountIDs) +} + +func TestAdminService_UpdateCompositeGroupCopiesAccountsFromConcreteGroups(t *testing.T) { + var clearedGroupID int64 + var copiedFrom []int64 + var boundGroupID int64 + var boundAccountIDs []int64 + groupRepo := &groupRepoStubForAdmin{ + getByIDByID: map[int64]*Group{ + 10: {ID: 10, Platform: PlatformOpenAI}, + 20: {ID: 20, Platform: PlatformGrok}, + 99: {ID: 99, Platform: PlatformComposite, RateMultiplier: 1, SubscriptionType: SubscriptionTypeStandard}, + }, + deleteAccountGroupsByGroupIDFn: func(groupID int64) (int64, error) { + clearedGroupID = groupID + return 2, nil + }, + getAccountIDsByGroupIDsFn: func(groupIDs []int64) ([]int64, error) { + copiedFrom = append([]int64{}, groupIDs...) + return []int64{301, 302}, nil + }, + bindAccountsToGroupFn: func(groupID int64, accountIDs []int64) error { + boundGroupID = groupID + boundAccountIDs = append([]int64{}, accountIDs...) + return nil + }, + } + svc := &adminServiceImpl{groupRepo: groupRepo} + + group, err := svc.UpdateGroup(context.Background(), 99, &UpdateGroupInput{ + CopyAccountsFromGroupIDs: []int64{10, 20}, + }) + + require.NoError(t, err) + require.Equal(t, PlatformComposite, group.Platform) + require.Equal(t, int64(99), clearedGroupID) + require.ElementsMatch(t, []int64{10, 20}, copiedFrom) + require.Equal(t, int64(99), boundGroupID) + require.ElementsMatch(t, []int64{301, 302}, boundAccountIDs) +} + +func TestAdminService_CreateAccountAllowsCompositeGroupAssignment(t *testing.T) { + accountRepo := &accountRepoStubForBulkUpdate{createID: 7} + groupRepo := &groupRepoStubForAdmin{ + getByIDByID: map[int64]*Group{ + 99: {ID: 99, Platform: PlatformComposite}, + }, + } + svc := &adminServiceImpl{accountRepo: accountRepo, groupRepo: groupRepo} + + account, err := svc.CreateAccount(context.Background(), &CreateAccountInput{ + Name: "OpenAI account", + Platform: PlatformOpenAI, + Type: AccountTypeAPIKey, + Concurrency: 1, + GroupIDs: []int64{99}, + SkipDefaultGroupBind: true, + SkipMixedChannelCheck: true, + }) + + require.NoError(t, err) + require.Equal(t, int64(7), account.ID) + require.Equal(t, PlatformOpenAI, accountRepo.createAccount.Platform) + require.ElementsMatch(t, []int64{99}, accountRepo.bindGroupsByAccount[7]) +} + +func TestAdminService_UpdateAccountAllowsCompositeGroupAssignment(t *testing.T) { + accountRepo := &accountRepoStubForBulkUpdate{ + getByIDAccounts: map[int64]*Account{ + 7: {ID: 7, Platform: PlatformGemini, Type: AccountTypeAPIKey, Status: StatusActive, Extra: map[string]any{}}, + }, + } + groupRepo := &groupRepoStubForAdmin{ + getByIDByID: map[int64]*Group{ + 99: {ID: 99, Platform: PlatformComposite}, + }, + } + svc := &adminServiceImpl{accountRepo: accountRepo, groupRepo: groupRepo} + groupIDs := []int64{99} + + account, err := svc.UpdateAccount(context.Background(), 7, &UpdateAccountInput{ + GroupIDs: &groupIDs, + SkipMixedChannelCheck: true, + }) + + require.NoError(t, err) + require.Equal(t, int64(7), account.ID) + require.Len(t, accountRepo.updatedAccounts, 1) + require.ElementsMatch(t, []int64{99}, accountRepo.bindGroupsByAccount[7]) +} + +func TestAdminService_CompositeModelsListCandidatesIncludeConcreteAccountMappings(t *testing.T) { + accountRepo := &accountRepoStubForCompositeModelsList{ + accounts: []Account{ + { + ID: 1, + Platform: PlatformOpenAI, + Credentials: map[string]any{ + "model_mapping": map[string]any{"gpt-custom": "gpt-5"}, + }, + }, + { + ID: 2, + Platform: PlatformGemini, + Credentials: map[string]any{ + "model_mapping": map[string]any{"gemini-custom": "gemini-2.5-flash"}, + }, + }, + }, + } + groupRepo := &groupRepoStubForAdmin{ + getByIDByID: map[int64]*Group{ + 99: {ID: 99, Platform: PlatformComposite}, + }, + } + svc := &adminServiceImpl{accountRepo: accountRepo, groupRepo: groupRepo} + + candidates, err := svc.GetGroupModelsListCandidates(context.Background(), 99, PlatformComposite) + + require.NoError(t, err) + require.Contains(t, candidates, "gpt-custom") + require.Contains(t, candidates, "gemini-custom") + require.Contains(t, candidates, "gpt-5.5") + require.Contains(t, candidates, "gemini-2.5-flash") +} diff --git a/backend/internal/service/admin_service_group_test.go b/backend/internal/service/admin_service_group_test.go index 5a6dccc115..5e3ec23fb8 100644 --- a/backend/internal/service/admin_service_group_test.go +++ b/backend/internal/service/admin_service_group_test.go @@ -17,10 +17,17 @@ func ptrString[T ~string](v T) *string { // groupRepoStubForAdmin 用于测试 AdminService 的 GroupRepository Stub type groupRepoStubForAdmin struct { - created *Group // 记录 Create 调用的参数 - updated *Group // 记录 Update 调用的参数 - getByID *Group // GetByID 返回值 - getErr error // GetByID 返回的错误 + created *Group // 记录 Create 调用的参数 + updated *Group // 记录 Update 调用的参数 + getByID *Group // GetByID 返回值 + getErr error // GetByID 返回的错误 + createID int64 + + getByIDByID map[int64]*Group + + deleteAccountGroupsByGroupIDFn func(groupID int64) (int64, error) + bindAccountsToGroupFn func(groupID int64, accountIDs []int64) error + getAccountIDsByGroupIDsFn func(groupIDs []int64) ([]int64, error) listWithFiltersCalls int listWithFiltersParams pagination.PaginationParams @@ -34,6 +41,9 @@ type groupRepoStubForAdmin struct { } func (s *groupRepoStubForAdmin) Create(_ context.Context, g *Group) error { + if s.createID > 0 { + g.ID = s.createID + } s.created = g return nil } @@ -43,17 +53,29 @@ func (s *groupRepoStubForAdmin) Update(_ context.Context, g *Group) error { return nil } -func (s *groupRepoStubForAdmin) GetByID(_ context.Context, _ int64) (*Group, error) { +func (s *groupRepoStubForAdmin) GetByID(_ context.Context, id int64) (*Group, error) { if s.getErr != nil { return nil, s.getErr } + if s.getByIDByID != nil { + if group, ok := s.getByIDByID[id]; ok { + return group, nil + } + return nil, ErrGroupNotFound + } return s.getByID, nil } -func (s *groupRepoStubForAdmin) GetByIDLite(_ context.Context, _ int64) (*Group, error) { +func (s *groupRepoStubForAdmin) GetByIDLite(_ context.Context, id int64) (*Group, error) { if s.getErr != nil { return nil, s.getErr } + if s.getByIDByID != nil { + if group, ok := s.getByIDByID[id]; ok { + return group, nil + } + return nil, ErrGroupNotFound + } return s.getByID, nil } @@ -109,15 +131,24 @@ func (s *groupRepoStubForAdmin) GetAccountCount(_ context.Context, _ int64) (int panic("unexpected GetAccountCount call") } -func (s *groupRepoStubForAdmin) DeleteAccountGroupsByGroupID(_ context.Context, _ int64) (int64, error) { +func (s *groupRepoStubForAdmin) DeleteAccountGroupsByGroupID(_ context.Context, groupID int64) (int64, error) { + if s.deleteAccountGroupsByGroupIDFn != nil { + return s.deleteAccountGroupsByGroupIDFn(groupID) + } panic("unexpected DeleteAccountGroupsByGroupID call") } -func (s *groupRepoStubForAdmin) BindAccountsToGroup(_ context.Context, _ int64, _ []int64) error { +func (s *groupRepoStubForAdmin) BindAccountsToGroup(_ context.Context, groupID int64, accountIDs []int64) error { + if s.bindAccountsToGroupFn != nil { + return s.bindAccountsToGroupFn(groupID, accountIDs) + } panic("unexpected BindAccountsToGroup call") } -func (s *groupRepoStubForAdmin) GetAccountIDsByGroupIDs(_ context.Context, _ []int64) ([]int64, error) { +func (s *groupRepoStubForAdmin) GetAccountIDsByGroupIDs(_ context.Context, groupIDs []int64) ([]int64, error) { + if s.getAccountIDsByGroupIDsFn != nil { + return s.getAccountIDsByGroupIDsFn(groupIDs) + } panic("unexpected GetAccountIDsByGroupIDs call") } diff --git a/backend/internal/service/channel_service.go b/backend/internal/service/channel_service.go index 25c7b587e3..a3a14893ab 100644 --- a/backend/internal/service/channel_service.go +++ b/backend/internal/service/channel_service.go @@ -8,6 +8,7 @@ import ( "sync/atomic" "time" + "github.com/Wei-Shaw/sub2api/internal/pkg/ctxkey" infraerrors "github.com/Wei-Shaw/sub2api/internal/pkg/errors" "github.com/Wei-Shaw/sub2api/internal/pkg/pagination" "github.com/tidwall/gjson" @@ -332,16 +333,38 @@ func populateChannelCache(channels []Channel, groupPlatforms map[int64]string) * // invalidateCache 使缓存失效,让下次读取时自然重建 // isPlatformPricingMatch 判断定价条目的平台是否匹配分组平台。 -// 各平台(antigravity / anthropic / gemini / openai)严格独立,不跨平台匹配。 +// Concrete platforms stay isolated; composite groups may carry concrete-provider +// pricing rows that are selected by the request's resolved target platform. func isPlatformPricingMatch(groupPlatform, pricingPlatform string) bool { + if groupPlatform == PlatformComposite { + return isConcreteRequestPlatform(pricingPlatform) + } return groupPlatform == pricingPlatform } // matchingPlatforms 返回分组平台对应的可匹配平台列表。 -// 各平台严格独立,只返回自身。 +// Concrete platforms return themselves; composite is a configuration-time +// 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{groupPlatform} } + +func channelLookupPlatform(ctx context.Context, groupPlatform string) string { + if ctx != nil { + if forcePlatform, ok := ctx.Value(ctxkey.ForcePlatform).(string); ok && strings.TrimSpace(forcePlatform) != "" { + return strings.TrimSpace(forcePlatform) + } + if groupPlatform == PlatformComposite { + if platform, ok := ResolvedTargetPlatformFromContext(ctx); ok { + return platform + } + } + } + return groupPlatform +} func (s *ChannelService) invalidateCache() { s.cache.Store((*channelCache)(nil)) s.cacheSF.Forget("channel_cache") @@ -456,7 +479,7 @@ func (s *ChannelService) lookupGroupChannel(ctx context.Context, groupID int64) return &channelLookup{ cache: cache, channel: ch, - platform: cache.groupPlatform[groupID], + platform: channelLookupPlatform(ctx, cache.groupPlatform[groupID]), }, nil } diff --git a/backend/internal/service/channel_service_test.go b/backend/internal/service/channel_service_test.go index 381b8c6c25..035b941dbd 100644 --- a/backend/internal/service/channel_service_test.go +++ b/backend/internal/service/channel_service_test.go @@ -1970,6 +1970,8 @@ func TestIsPlatformPricingMatch(t *testing.T) { {"gemini matches gemini", PlatformGemini, PlatformGemini, true}, {"gemini does NOT match antigravity", PlatformGemini, PlatformAntigravity, false}, {"gemini does NOT match anthropic", PlatformGemini, PlatformAnthropic, false}, + {"composite matches openai pricing", PlatformComposite, PlatformOpenAI, true}, + {"composite matches gemini pricing", PlatformComposite, PlatformGemini, true}, {"empty string matches nothing", "", PlatformAnthropic, false}, {"empty string matches empty", "", "", true}, } @@ -1995,6 +1997,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}}, } for _, tt := range tests { @@ -2005,6 +2008,43 @@ func TestMatchingPlatforms(t *testing.T) { } } +func TestCompositeChannelLookupUsesResolvedTargetPlatform(t *testing.T) { + channel := Channel{ + ID: 1, + Status: StatusActive, + GroupIDs: []int64{99}, + ModelPricing: []ChannelModelPricing{ + {Platform: PlatformOpenAI, Models: []string{"gpt-*"}}, + {Platform: PlatformAnthropic, Models: []string{"claude-*"}}, + }, + ModelMapping: map[string]map[string]string{ + PlatformOpenAI: { + "gpt-5": "gpt-5-mini", + }, + PlatformAnthropic: { + "claude-*": "claude-sonnet-4-5", + }, + }, + } + cache := populateChannelCache([]Channel{channel}, map[int64]string{99: PlatformComposite}) + svc := &ChannelService{} + svc.cache.Store(cache) + + openAICtx := WithResolvedTargetPlatform(context.Background(), PlatformOpenAI) + require.NotNil(t, svc.GetChannelModelPricing(openAICtx, 99, "gpt-5")) + require.Nil(t, svc.GetChannelModelPricing(openAICtx, 99, "claude-sonnet-4-5")) + openAIResult := svc.ResolveChannelMapping(openAICtx, 99, "gpt-5") + require.True(t, openAIResult.Mapped) + require.Equal(t, "gpt-5-mini", openAIResult.MappedModel) + + anthropicCtx := WithResolvedTargetPlatform(context.Background(), PlatformAnthropic) + require.NotNil(t, svc.GetChannelModelPricing(anthropicCtx, 99, "claude-sonnet-4-5")) + require.Nil(t, svc.GetChannelModelPricing(anthropicCtx, 99, "gpt-5")) + anthropicResult := svc.ResolveChannelMapping(anthropicCtx, 99, "claude-3-5-sonnet") + require.True(t, anthropicResult.Mapped) + require.Equal(t, "claude-sonnet-4-5", anthropicResult.MappedModel) +} + // =========================================================================== // 9. Antigravity platform isolation — no cross-platform pricing leakage // =========================================================================== diff --git a/docs/COMPOSITE_GROUPS.md b/docs/COMPOSITE_GROUPS.md new file mode 100644 index 0000000000..f8397c0faf --- /dev/null +++ b/docs/COMPOSITE_GROUPS.md @@ -0,0 +1,49 @@ +# Composite Groups + +Composite groups are an admin routing layer for API keys that should choose a +concrete provider from the requested model instead of binding the key to a +single provider group. + +## Supported Providers + +Composite groups can route to these concrete account platforms: + +- Anthropic +- Gemini +- OpenAI +- Antigravity +- Grok + +The selected concrete platform is used for account selection, user platform +quota checks, post-usage billing, ops error platform attribution, channel +mapping/pricing lookup, and platform usage reporting. + +## Model Detection + +Composite routing detects common public model IDs and provider-prefixed IDs: + +- `claude-*` and `anthropic/claude-*` route to Anthropic. +- `gemini-*` and `google/gemini-*` route to Gemini. +- `gpt-*`, `o*`, `codex-*`, `text-embedding-*`, `dall-e-*`, and + `openai/*` route to OpenAI. +- `grok-*` and `xai/grok-*` route to Grok. + +Unknown or ambiguous model names fail closed with a client error instead of +guessing a provider. + +## Admin Workflows + +- Admins can create a group with platform `composite`. +- Composite groups can copy accounts from concrete provider groups. +- Concrete provider accounts can be assigned directly to composite groups from + account create/edit and bulk account workflows. +- Channel configuration exposes composite groups in concrete provider sections. + The channel `group_ids` payload is still flat; provider-specific model + mapping and pricing remain keyed by concrete platform. + +## Limits + +Composite groups are not a full OpenRouter-compatible model registry. They do +not add a provider/model mapping database, per-model admin routing overrides, or +arbitrary third-party provider prefixes. Add those explicitly before relying on +custom model IDs that cannot be detected from their names. diff --git a/frontend/src/views/admin/ChannelsView.vue b/frontend/src/views/admin/ChannelsView.vue index defe470159..befc356f2b 100644 --- a/frontend/src/views/admin/ChannelsView.vue +++ b/frontend/src/views/admin/ChannelsView.vue @@ -799,7 +799,7 @@ function togglePlatform(platform: GroupPlatform) { } function getGroupsForPlatform(platform: GroupPlatform): AdminGroup[] { - return allGroups.value.filter(g => g.platform === platform) + return allGroups.value.filter(g => g.platform === platform || g.platform === 'composite') } // ── Group helpers ── @@ -1117,6 +1117,7 @@ function formToAPI(): { group_ids: number[], model_pricing: ChannelModelPricing[ }) } } + const uniqueGroupIds = Array.from(new Set(group_ids)) // Collect web_search_emulation (only anthropic platform supports it) // Always write the key so that disabling in the UI correctly sets platform to false, @@ -1160,7 +1161,7 @@ function formToAPI(): { group_ids: number[], model_pricing: ChannelModelPricing[ delete featuresConfig.bedrock_cc_compat } - return { group_ids, model_pricing, model_mapping, features_config: featuresConfig } + return { group_ids: uniqueGroupIds, model_pricing, model_mapping, features_config: featuresConfig } } function apiToForm(channel: Channel): PlatformSection[] { @@ -1174,7 +1175,11 @@ function apiToForm(channel: Channel): PlatformSection[] { const activePlatforms = new Set() for (const gid of channel.group_ids || []) { const p = groupPlatformMap.get(gid) - if (p) activePlatforms.add(p) + if (p === 'composite') { + platformOrder.forEach(platform => activePlatforms.add(platform)) + } else if (p) { + activePlatforms.add(p) + } } for (const p of channel.model_pricing || []) { if (p.platform) activePlatforms.add(p.platform as GroupPlatform) @@ -1188,7 +1193,10 @@ function apiToForm(channel: Channel): PlatformSection[] { for (const platform of platformOrder) { if (!activePlatforms.has(platform)) continue - const groupIds = (channel.group_ids || []).filter(gid => groupPlatformMap.get(gid) === platform) + const groupIds = (channel.group_ids || []).filter(gid => { + const groupPlatform = groupPlatformMap.get(gid) + return groupPlatform === platform || groupPlatform === 'composite' + }) const mapping = (channel.model_mapping || {})[platform] || {} const pricing = (channel.model_pricing || []) .filter(p => (p.platform || 'anthropic') === platform) @@ -1364,7 +1372,7 @@ function distributeRulesToPlatforms(apiRules: AccountStatsPricingRule[]) { const platforms = new Set() for (const gid of apiRule.group_ids || []) { const p = groupPlatformMap.get(gid) - if (p) platforms.add(p) + if (p && p !== 'composite') platforms.add(p) } // If pricing has a platform field, use that as fallback if (platforms.size === 0 && apiRule.pricing?.length > 0) {