mirror of
https://github.com/Wei-Shaw/sub2api.git
synced 2026-10-07 13:28:39 +08:00
Harden composite group product surfaces
This commit is contained in:
@@ -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)
|
||||
|
||||
@@ -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{
|
||||
|
||||
@@ -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")
|
||||
}
|
||||
|
||||
|
||||
@@ -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
|
||||
// ===========================================================================
|
||||
|
||||
@@ -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.
|
||||
@@ -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) {
|
||||
|
||||
Reference in New Issue
Block a user