Harden composite group product surfaces

This commit is contained in:
Heatherm Huang
2026-07-23 09:19:25 +08:00
parent ebc1028771
commit c8d1e2e16f
20 changed files with 570 additions and 45 deletions
+1
View File
@@ -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
@@ -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 {
@@ -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"))
}
@@ -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)))
@@ -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)
@@ -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)
@@ -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)
@@ -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)
+6 -3
View File
@@ -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
}
@@ -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))
}
@@ -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")
}
@@ -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{
+21 -2
View File
@@ -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)
@@ -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
}
@@ -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")
}
@@ -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")
}
+26 -3
View File
@@ -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
}
@@ -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
// ===========================================================================
+49
View File
@@ -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.
+13 -5
View File
@@ -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<GroupPlatform>()
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<GroupPlatform>()
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) {