|
- 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)
- }
|