Não pode escolher mais do que 25 tópicos Os tópicos devem começar com uma letra ou um número, podem incluir traços ('-') e podem ter até 35 caracteres.
 
 
 

242 linhas
7.0 KiB

  1. package model
  2. import (
  3. "sync"
  4. "testing"
  5. "github.com/glebarez/sqlite"
  6. "github.com/stretchr/testify/assert"
  7. "github.com/stretchr/testify/require"
  8. "gorm.io/gorm"
  9. "gorm.io/gorm/logger"
  10. )
  11. func setupUserAssetChannelDB(t *testing.T) *gorm.DB {
  12. t.Helper()
  13. db, err := gorm.Open(sqlite.Open("file:user_asset_channels?mode=memory&cache=shared"), &gorm.Config{
  14. Logger: logger.Default.LogMode(logger.Silent),
  15. })
  16. require.NoError(t, err)
  17. sqlDB, err := db.DB()
  18. require.NoError(t, err)
  19. sqlDB.SetMaxOpenConns(1)
  20. origDB := DB
  21. origGroupCol := commonGroupCol
  22. DB = db
  23. commonGroupCol = "`group`"
  24. require.NoError(t, db.AutoMigrate(&UserAssetChannel{}))
  25. t.Cleanup(func() {
  26. DB = origDB
  27. commonGroupCol = origGroupCol
  28. require.NoError(t, sqlDB.Close())
  29. })
  30. return db
  31. }
  32. func TestGetUserAssetChannel_NotFound(t *testing.T) {
  33. setupUserAssetChannelDB(t)
  34. binding, err := GetUserAssetChannel(1, 2, "default")
  35. require.NoError(t, err)
  36. assert.Nil(t, binding)
  37. }
  38. func TestBindUserAssetChannel_CreateAndUpdate(t *testing.T) {
  39. db := setupUserAssetChannelDB(t)
  40. require.NoError(t, BindUserAssetChannel(1, 2, "default", 100))
  41. binding, err := GetUserAssetChannel(1, 2, "default")
  42. require.NoError(t, err)
  43. require.NotNil(t, binding)
  44. assert.Equal(t, 100, binding.ChannelId)
  45. assert.NotZero(t, binding.CreatedAt)
  46. assert.NotZero(t, binding.UpdatedAt)
  47. require.NoError(t, BindUserAssetChannel(1, 2, "default", 200))
  48. binding, err = GetUserAssetChannel(1, 2, "default")
  49. require.NoError(t, err)
  50. require.NotNil(t, binding)
  51. assert.Equal(t, 200, binding.ChannelId)
  52. var count int64
  53. require.NoError(t, db.Model(&UserAssetChannel{}).Count(&count).Error)
  54. assert.Equal(t, int64(1), count)
  55. }
  56. func TestGetUserAssetChannelsByTypesSortsLatestFirst(t *testing.T) {
  57. db := setupUserAssetChannelDB(t)
  58. require.NoError(t, BindUserAssetChannel(1, 2, "default", 100))
  59. require.NoError(t, BindUserAssetChannel(1, 3, "default", 200))
  60. require.NoError(t, BindUserAssetChannel(1, 4, "default", 300))
  61. require.NoError(t, db.Model(&UserAssetChannel{}).
  62. Where("user_id = ? AND channel_type = ?", 1, 2).
  63. Update("updated_at", int64(100)).Error)
  64. require.NoError(t, db.Model(&UserAssetChannel{}).
  65. Where("user_id = ? AND channel_type = ?", 1, 3).
  66. Update("updated_at", int64(200)).Error)
  67. bindings, err := GetUserAssetChannelsByTypes(1, []int{2, 3}, "default")
  68. require.NoError(t, err)
  69. require.Len(t, bindings, 2)
  70. assert.Equal(t, 3, bindings[0].ChannelType)
  71. assert.Equal(t, 200, bindings[0].ChannelId)
  72. assert.Equal(t, 2, bindings[1].ChannelType)
  73. }
  74. func TestGetUserAssetChannelsReturnsAllTypesSortedLatestFirst(t *testing.T) {
  75. db := setupUserAssetChannelDB(t)
  76. require.NoError(t, BindUserAssetChannel(1, 2, "default", 100))
  77. require.NoError(t, BindUserAssetChannel(1, 3, "default", 200))
  78. require.NoError(t, BindUserAssetChannel(1, 4, "vip", 300))
  79. require.NoError(t, db.Model(&UserAssetChannel{}).
  80. Where("user_id = ? AND channel_type = ?", 1, 2).
  81. Update("updated_at", int64(100)).Error)
  82. require.NoError(t, db.Model(&UserAssetChannel{}).
  83. Where("user_id = ? AND channel_type = ?", 1, 3).
  84. Update("updated_at", int64(200)).Error)
  85. bindings, err := GetUserAssetChannels(1, "default")
  86. require.NoError(t, err)
  87. require.Len(t, bindings, 2)
  88. assert.Equal(t, 3, bindings[0].ChannelType)
  89. assert.Equal(t, 2, bindings[1].ChannelType)
  90. }
  91. func TestBindUserAssetChannelWithTxUpserts(t *testing.T) {
  92. db := setupUserAssetChannelDB(t)
  93. require.NoError(t, BindUserAssetChannel(1, 2, "default", 100))
  94. require.NoError(t, BindUserAssetChannel(1, 2, "default", 200))
  95. binding, err := GetUserAssetChannel(1, 2, "default")
  96. require.NoError(t, err)
  97. require.NotNil(t, binding)
  98. assert.Equal(t, 200, binding.ChannelId)
  99. var count int64
  100. require.NoError(t, db.Model(&UserAssetChannel{}).Count(&count).Error)
  101. assert.Equal(t, int64(1), count)
  102. }
  103. func TestDeleteUserAssetChannelsByTypesWithTxDeletesOnlyRequestedRows(t *testing.T) {
  104. db := setupUserAssetChannelDB(t)
  105. require.NoError(t, BindUserAssetChannel(1, 2, "default", 100))
  106. require.NoError(t, BindUserAssetChannel(1, 3, "default", 200))
  107. require.NoError(t, BindUserAssetChannel(1, 4, "default", 300))
  108. require.NoError(t, BindUserAssetChannel(1, 2, "vip", 400))
  109. require.NoError(t, BindUserAssetChannel(2, 2, "default", 500))
  110. require.NoError(t, db.Transaction(func(tx *gorm.DB) error {
  111. return DeleteUserAssetChannelsByTypesWithTx(tx, 1, []int{2, 3}, "default")
  112. }))
  113. binding, err := GetUserAssetChannel(1, 2, "default")
  114. require.NoError(t, err)
  115. assert.Nil(t, binding)
  116. binding, err = GetUserAssetChannel(1, 3, "default")
  117. require.NoError(t, err)
  118. assert.Nil(t, binding)
  119. binding, err = GetUserAssetChannel(1, 4, "default")
  120. require.NoError(t, err)
  121. require.NotNil(t, binding)
  122. assert.Equal(t, 300, binding.ChannelId)
  123. binding, err = GetUserAssetChannel(1, 2, "vip")
  124. require.NoError(t, err)
  125. require.NotNil(t, binding)
  126. assert.Equal(t, 400, binding.ChannelId)
  127. binding, err = GetUserAssetChannel(2, 2, "default")
  128. require.NoError(t, err)
  129. require.NotNil(t, binding)
  130. assert.Equal(t, 500, binding.ChannelId)
  131. }
  132. func TestBindUserAssetChannel_IsolatesUserTypeGroup(t *testing.T) {
  133. setupUserAssetChannelDB(t)
  134. require.NoError(t, BindUserAssetChannel(1, 2, "default", 100))
  135. require.NoError(t, BindUserAssetChannel(2, 2, "default", 200))
  136. require.NoError(t, BindUserAssetChannel(1, 3, "default", 300))
  137. require.NoError(t, BindUserAssetChannel(1, 2, "vip", 400))
  138. cases := []struct {
  139. userId int
  140. channelType int
  141. group string
  142. channelId int
  143. }{
  144. {1, 2, "default", 100},
  145. {2, 2, "default", 200},
  146. {1, 3, "default", 300},
  147. {1, 2, "vip", 400},
  148. }
  149. for _, tc := range cases {
  150. binding, err := GetUserAssetChannel(tc.userId, tc.channelType, tc.group)
  151. require.NoError(t, err)
  152. require.NotNil(t, binding)
  153. assert.Equal(t, tc.channelId, binding.ChannelId)
  154. }
  155. }
  156. func TestBindUserAssetChannel_ConcurrentUpsert(t *testing.T) {
  157. db := setupUserAssetChannelDB(t)
  158. const workers = 20
  159. var wg sync.WaitGroup
  160. errCh := make(chan error, workers)
  161. expected := make(map[int]bool, workers)
  162. for i := 0; i < workers; i++ {
  163. channelId := 1000 + i
  164. expected[channelId] = true
  165. wg.Add(1)
  166. go func() {
  167. defer wg.Done()
  168. errCh <- BindUserAssetChannel(1, 2, "default", channelId)
  169. }()
  170. }
  171. wg.Wait()
  172. close(errCh)
  173. for err := range errCh {
  174. require.NoError(t, err)
  175. }
  176. var count int64
  177. require.NoError(t, db.Model(&UserAssetChannel{}).Count(&count).Error)
  178. assert.Equal(t, int64(1), count)
  179. binding, err := GetUserAssetChannel(1, 2, "default")
  180. require.NoError(t, err)
  181. require.NotNil(t, binding)
  182. assert.True(t, expected[binding.ChannelId])
  183. }
  184. func TestUnbindUserAssetChannel(t *testing.T) {
  185. setupUserAssetChannelDB(t)
  186. require.NoError(t, BindUserAssetChannel(1, 2, "default", 100))
  187. require.NoError(t, BindUserAssetChannel(1, 2, "vip", 200))
  188. require.NoError(t, UnbindUserAssetChannel(1, 2, "default"))
  189. binding, err := GetUserAssetChannel(1, 2, "default")
  190. require.NoError(t, err)
  191. assert.Nil(t, binding)
  192. binding, err = GetUserAssetChannel(1, 2, "vip")
  193. require.NoError(t, err)
  194. require.NotNil(t, binding)
  195. assert.Equal(t, 200, binding.ChannelId)
  196. }