You can not select more than 25 topics Topics must start with a letter or number, can include dashes ('-') and can be up to 35 characters long.
 
 
 

160 line
5.7 KiB

  1. package service
  2. import (
  3. "context"
  4. "net/http"
  5. "sync"
  6. "testing"
  7. "github.com/QuantumNous/new-api/common"
  8. "github.com/QuantumNous/new-api/constant"
  9. "github.com/QuantumNous/new-api/model"
  10. "github.com/glebarez/sqlite"
  11. "github.com/stretchr/testify/assert"
  12. "github.com/stretchr/testify/require"
  13. "gorm.io/gorm"
  14. "gorm.io/gorm/logger"
  15. )
  16. func setupChinaMobileUserAssetGroupDB(t *testing.T) *gorm.DB {
  17. t.Helper()
  18. db, err := gorm.Open(sqlite.Open("file:chinamobile_user_asset_groups?mode=memory&cache=shared"), &gorm.Config{Logger: logger.Default.LogMode(logger.Silent)})
  19. require.NoError(t, err)
  20. sqlDB, err := db.DB()
  21. require.NoError(t, err)
  22. sqlDB.SetMaxOpenConns(1)
  23. origDB := model.DB
  24. model.DB = db
  25. require.NoError(t, db.AutoMigrate(&model.UserAssetGroup{}))
  26. t.Cleanup(func() {
  27. model.DB = origDB
  28. require.NoError(t, sqlDB.Close())
  29. })
  30. return db
  31. }
  32. func chinaMobileAssetAction(t *testing.T, action string) AssetActionSpec {
  33. t.Helper()
  34. spec, ok := ParseAssetAction(action)
  35. require.True(t, ok)
  36. return spec
  37. }
  38. func TestGetOrCreateChinaMobileUserAssetGroupCreatesThenReuses(t *testing.T) {
  39. setupChinaMobileUserAssetGroupDB(t)
  40. channel := &model.Channel{Id: 101, Type: constant.ChannelTypeChinaMobileSeedance}
  41. createCalls := 0
  42. creator := func(context.Context, int, *model.Channel) (string, *AssetError) {
  43. createCalls++
  44. return "group-1", nil
  45. }
  46. groupID, assetErr := GetOrCreateChinaMobileUserAssetGroup(context.Background(), 10, channel, creator)
  47. require.Nil(t, assetErr)
  48. assert.Equal(t, "group-1", groupID)
  49. again, assetErr := GetOrCreateChinaMobileUserAssetGroup(context.Background(), 10, channel, creator)
  50. require.Nil(t, assetErr)
  51. assert.Equal(t, "group-1", again)
  52. assert.Equal(t, 1, createCalls)
  53. }
  54. func TestGetOrCreateChinaMobileUserAssetGroupConcurrentRequestsReuseBinding(t *testing.T) {
  55. setupChinaMobileUserAssetGroupDB(t)
  56. channel := &model.Channel{Id: 102, Type: constant.ChannelTypeChinaMobileSeedance}
  57. creator := func(context.Context, int, *model.Channel) (string, *AssetError) {
  58. return "group-concurrent", nil
  59. }
  60. const workers = 8
  61. results := make(chan string, workers)
  62. errs := make(chan *AssetError, workers)
  63. var wg sync.WaitGroup
  64. for i := 0; i < workers; i++ {
  65. wg.Add(1)
  66. go func() {
  67. defer wg.Done()
  68. groupID, assetErr := GetOrCreateChinaMobileUserAssetGroup(context.Background(), 10, channel, creator)
  69. results <- groupID
  70. errs <- assetErr
  71. }()
  72. }
  73. wg.Wait()
  74. close(results)
  75. close(errs)
  76. for assetErr := range errs {
  77. require.Nil(t, assetErr)
  78. }
  79. for groupID := range results {
  80. assert.Equal(t, "group-concurrent", groupID)
  81. }
  82. binding, err := model.GetUserAssetGroup(10, 102)
  83. require.NoError(t, err)
  84. require.NotNil(t, binding)
  85. assert.Equal(t, "group-concurrent", binding.GroupId)
  86. }
  87. func TestScopeChinaMobileAssetRequestOverwritesClientGroup(t *testing.T) {
  88. create := AssetRequest{Action: chinaMobileAssetAction(t, "CreateAsset"), Body: map[string]any{"GroupId": "forged"}}
  89. ScopeChinaMobileAssetRequest(&create, "owned")
  90. assert.Equal(t, "owned", create.Body["GroupId"])
  91. list := AssetRequest{Action: chinaMobileAssetAction(t, "ListAssets"), Body: map[string]any{"Filter": map[string]any{"GroupIds": []any{"forged"}}}}
  92. ScopeChinaMobileAssetRequest(&list, "owned")
  93. assert.Equal(t, []string{"owned"}, list.Body["Filter"].(map[string]any)["GroupIds"])
  94. }
  95. func TestRequireChinaMobileAssetOwnershipHidesMismatchedGroup(t *testing.T) {
  96. adapter := &recordingAssetAdapter{getAssetGroupID: "other"}
  97. request := AssetRequest{Action: chinaMobileAssetAction(t, "DeleteAsset"), Body: map[string]any{"Id": "asset-1"}}
  98. assetErr := RequireChinaMobileAssetOwnership(context.Background(), adapter, &model.Channel{}, request, "owned")
  99. require.NotNil(t, assetErr)
  100. assert.Equal(t, AssetErrorNotFound, assetErr.Type)
  101. assert.Equal(t, http.StatusNotFound, assetErr.HTTPStatus)
  102. assert.Equal(t, AssetOperationAssetGet, adapter.requests[0].Action.Operation)
  103. }
  104. type recordingAssetAdapter struct {
  105. requests []AssetRequest
  106. getAssetGroupID string
  107. }
  108. func (a *recordingAssetAdapter) Name() string { return "recording" }
  109. func (a *recordingAssetAdapter) Supports(AssetOperation) bool { return true }
  110. func (a *recordingAssetAdapter) DoAssetRequest(_ context.Context, _ *model.Channel, req AssetRequest) (*AssetUpstreamResponse, *AssetError) {
  111. a.requests = append(a.requests, req)
  112. if req.Action.Operation == AssetOperationAssetGroupCreate {
  113. body, err := BuildAssetSuccessResponse(req.Action.Action, req.Version, map[string]any{"GroupId": "group-created"})
  114. if err != nil {
  115. return nil, newAssetError(AssetErrorServer, err.Error(), http.StatusInternalServerError)
  116. }
  117. return &AssetUpstreamResponse{StatusCode: http.StatusOK, Body: body}, nil
  118. }
  119. if req.Action.Operation != AssetOperationAssetGet {
  120. return &AssetUpstreamResponse{StatusCode: http.StatusOK}, nil
  121. }
  122. body, err := BuildAssetSuccessResponse(req.Action.Action, req.Version, map[string]any{"Id": "asset-1", "GroupId": a.getAssetGroupID})
  123. if err != nil {
  124. return nil, newAssetError(AssetErrorServer, err.Error(), http.StatusInternalServerError)
  125. }
  126. return &AssetUpstreamResponse{StatusCode: http.StatusOK, Body: body}, nil
  127. }
  128. func TestCreateChinaMobileUserAssetGroupUsesInternalGroupAction(t *testing.T) {
  129. adapter := &recordingAssetAdapter{}
  130. groupID, assetErr := CreateChinaMobileUserAssetGroup(context.Background(), 10, adapter, &model.Channel{Id: 101})
  131. require.Nil(t, assetErr)
  132. assert.Equal(t, "group-created", groupID)
  133. assert.Equal(t, AssetOperationAssetGroupCreate, adapter.requests[0].Action.Operation)
  134. assert.Equal(t, "AIGC", adapter.requests[0].Body["GroupType"])
  135. assert.Equal(t, "new-api-user-10-channel-101", adapter.requests[0].Body["Name"])
  136. // Ensure the normalized response parser remains wired through common JSON helpers.
  137. _, err := common.Marshal(adapter.requests[0].Body)
  138. require.NoError(t, err)
  139. }