package service import ( "context" "net/http" "sync" "testing" "github.com/QuantumNous/new-api/common" "github.com/QuantumNous/new-api/constant" "github.com/QuantumNous/new-api/model" "github.com/glebarez/sqlite" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" "gorm.io/gorm" "gorm.io/gorm/logger" ) func setupChinaMobileUserAssetGroupDB(t *testing.T) *gorm.DB { t.Helper() db, err := gorm.Open(sqlite.Open("file:chinamobile_user_asset_groups?mode=memory&cache=shared"), &gorm.Config{Logger: logger.Default.LogMode(logger.Silent)}) require.NoError(t, err) sqlDB, err := db.DB() require.NoError(t, err) sqlDB.SetMaxOpenConns(1) origDB := model.DB model.DB = db require.NoError(t, db.AutoMigrate(&model.UserAssetGroup{})) t.Cleanup(func() { model.DB = origDB require.NoError(t, sqlDB.Close()) }) return db } func chinaMobileAssetAction(t *testing.T, action string) AssetActionSpec { t.Helper() spec, ok := ParseAssetAction(action) require.True(t, ok) return spec } func TestGetOrCreateChinaMobileUserAssetGroupCreatesThenReuses(t *testing.T) { setupChinaMobileUserAssetGroupDB(t) channel := &model.Channel{Id: 101, Type: constant.ChannelTypeChinaMobileSeedance} createCalls := 0 creator := func(context.Context, int, *model.Channel) (string, *AssetError) { createCalls++ return "group-1", nil } groupID, assetErr := GetOrCreateChinaMobileUserAssetGroup(context.Background(), 10, channel, creator) require.Nil(t, assetErr) assert.Equal(t, "group-1", groupID) again, assetErr := GetOrCreateChinaMobileUserAssetGroup(context.Background(), 10, channel, creator) require.Nil(t, assetErr) assert.Equal(t, "group-1", again) assert.Equal(t, 1, createCalls) } func TestGetOrCreateChinaMobileUserAssetGroupConcurrentRequestsReuseBinding(t *testing.T) { setupChinaMobileUserAssetGroupDB(t) channel := &model.Channel{Id: 102, Type: constant.ChannelTypeChinaMobileSeedance} creator := func(context.Context, int, *model.Channel) (string, *AssetError) { return "group-concurrent", nil } const workers = 8 results := make(chan string, workers) errs := make(chan *AssetError, workers) var wg sync.WaitGroup for i := 0; i < workers; i++ { wg.Add(1) go func() { defer wg.Done() groupID, assetErr := GetOrCreateChinaMobileUserAssetGroup(context.Background(), 10, channel, creator) results <- groupID errs <- assetErr }() } wg.Wait() close(results) close(errs) for assetErr := range errs { require.Nil(t, assetErr) } for groupID := range results { assert.Equal(t, "group-concurrent", groupID) } binding, err := model.GetUserAssetGroup(10, 102) require.NoError(t, err) require.NotNil(t, binding) assert.Equal(t, "group-concurrent", binding.GroupId) } func TestScopeChinaMobileAssetRequestOverwritesClientGroup(t *testing.T) { create := AssetRequest{Action: chinaMobileAssetAction(t, "CreateAsset"), Body: map[string]any{"GroupId": "forged"}} ScopeChinaMobileAssetRequest(&create, "owned") assert.Equal(t, "owned", create.Body["GroupId"]) list := AssetRequest{Action: chinaMobileAssetAction(t, "ListAssets"), Body: map[string]any{"Filter": map[string]any{"GroupIds": []any{"forged"}}}} ScopeChinaMobileAssetRequest(&list, "owned") assert.Equal(t, []string{"owned"}, list.Body["Filter"].(map[string]any)["GroupIds"]) } func TestRequireChinaMobileAssetOwnershipHidesMismatchedGroup(t *testing.T) { adapter := &recordingAssetAdapter{getAssetGroupID: "other"} request := AssetRequest{Action: chinaMobileAssetAction(t, "DeleteAsset"), Body: map[string]any{"Id": "asset-1"}} assetErr := RequireChinaMobileAssetOwnership(context.Background(), adapter, &model.Channel{}, request, "owned") require.NotNil(t, assetErr) assert.Equal(t, AssetErrorNotFound, assetErr.Type) assert.Equal(t, http.StatusNotFound, assetErr.HTTPStatus) assert.Equal(t, AssetOperationAssetGet, adapter.requests[0].Action.Operation) } type recordingAssetAdapter struct { requests []AssetRequest getAssetGroupID string } func (a *recordingAssetAdapter) Name() string { return "recording" } func (a *recordingAssetAdapter) Supports(AssetOperation) bool { return true } func (a *recordingAssetAdapter) DoAssetRequest(_ context.Context, _ *model.Channel, req AssetRequest) (*AssetUpstreamResponse, *AssetError) { a.requests = append(a.requests, req) if req.Action.Operation == AssetOperationAssetGroupCreate { body, err := BuildAssetSuccessResponse(req.Action.Action, req.Version, map[string]any{"GroupId": "group-created"}) if err != nil { return nil, newAssetError(AssetErrorServer, err.Error(), http.StatusInternalServerError) } return &AssetUpstreamResponse{StatusCode: http.StatusOK, Body: body}, nil } if req.Action.Operation != AssetOperationAssetGet { return &AssetUpstreamResponse{StatusCode: http.StatusOK}, nil } body, err := BuildAssetSuccessResponse(req.Action.Action, req.Version, map[string]any{"Id": "asset-1", "GroupId": a.getAssetGroupID}) if err != nil { return nil, newAssetError(AssetErrorServer, err.Error(), http.StatusInternalServerError) } return &AssetUpstreamResponse{StatusCode: http.StatusOK, Body: body}, nil } func TestCreateChinaMobileUserAssetGroupUsesInternalGroupAction(t *testing.T) { adapter := &recordingAssetAdapter{} groupID, assetErr := CreateChinaMobileUserAssetGroup(context.Background(), 10, adapter, &model.Channel{Id: 101}) require.Nil(t, assetErr) assert.Equal(t, "group-created", groupID) assert.Equal(t, AssetOperationAssetGroupCreate, adapter.requests[0].Action.Operation) assert.Equal(t, "AIGC", adapter.requests[0].Body["GroupType"]) assert.Equal(t, "new-api-user-10-channel-101", adapter.requests[0].Body["Name"]) // Ensure the normalized response parser remains wired through common JSON helpers. _, err := common.Marshal(adapter.requests[0].Body) require.NoError(t, err) }