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.
 
 
 

244 lines
12 KiB

  1. package controller
  2. import (
  3. "bytes"
  4. "net/http"
  5. "net/http/httptest"
  6. "os"
  7. "strings"
  8. "testing"
  9. "github.com/QuantumNous/new-api/common"
  10. "github.com/QuantumNous/new-api/constant"
  11. "github.com/QuantumNous/new-api/model"
  12. "github.com/gin-gonic/gin"
  13. "github.com/stretchr/testify/assert"
  14. "github.com/stretchr/testify/require"
  15. "gorm.io/gorm"
  16. )
  17. func setupUserVideoChannelBindingDB(t *testing.T) *gorm.DB {
  18. t.Helper()
  19. originalDB := model.DB
  20. originalLogDB := model.LOG_DB
  21. originalCache := common.MemoryCacheEnabled
  22. originalRedisEnabled := common.RedisEnabled
  23. originalSQLitePath := common.SQLitePath
  24. originalIsMasterNode := common.IsMasterNode
  25. originalUsingSQLite := common.UsingSQLite
  26. originalUsingMySQL := common.UsingMySQL
  27. originalUsingPostgreSQL := common.UsingPostgreSQL
  28. originalSQLDSN, hadSQLDSN := os.LookupEnv("SQL_DSN")
  29. common.SQLitePath = "file:" + strings.NewReplacer("/", "_", " ", "_").Replace(t.Name()) + "?mode=memory&cache=shared"
  30. common.MemoryCacheEnabled = false
  31. common.RedisEnabled = false
  32. common.IsMasterNode = false
  33. common.UsingSQLite = false
  34. common.UsingMySQL = false
  35. common.UsingPostgreSQL = false
  36. require.NoError(t, os.Setenv("SQL_DSN", "local"))
  37. require.NoError(t, model.InitDB())
  38. db := model.DB
  39. model.LOG_DB = db
  40. sqlDB, err := db.DB()
  41. require.NoError(t, err)
  42. sqlDB.SetMaxOpenConns(1)
  43. require.NoError(t, db.AutoMigrate(&model.User{}, &model.Token{}, &model.Channel{}, &model.Ability{}, &model.UserAssetChannel{}, &model.Log{}))
  44. t.Cleanup(func() {
  45. model.DB = originalDB
  46. model.LOG_DB = originalLogDB
  47. common.MemoryCacheEnabled = originalCache
  48. common.RedisEnabled = originalRedisEnabled
  49. common.SQLitePath = originalSQLitePath
  50. common.IsMasterNode = originalIsMasterNode
  51. common.UsingSQLite = originalUsingSQLite
  52. common.UsingMySQL = originalUsingMySQL
  53. common.UsingPostgreSQL = originalUsingPostgreSQL
  54. if hadSQLDSN {
  55. _ = os.Setenv("SQL_DSN", originalSQLDSN)
  56. } else {
  57. _ = os.Unsetenv("SQL_DSN")
  58. }
  59. require.NoError(t, sqlDB.Close())
  60. })
  61. return db
  62. }
  63. func setupUserVideoChannelBindingRouter() *gin.Engine {
  64. gin.SetMode(gin.TestMode)
  65. router := gin.New()
  66. router.GET("/api/user/:id/video-channel-bindings", GetUserVideoChannelBindings)
  67. router.PUT("/api/user/:id/video-channel-bindings", SetUserVideoChannelBindings)
  68. return router
  69. }
  70. func createUserVideoBindingChannel(t *testing.T, db *gorm.DB, id, channelType int, group string, status int) {
  71. t.Helper()
  72. priority := int64(id)
  73. weight := uint(1)
  74. autoBan := 1
  75. require.NoError(t, db.Create(&model.Channel{Id: id, Type: channelType, Key: "channel-key", Status: status, Name: "channel", Group: group, Models: "model", Priority: &priority, Weight: &weight, AutoBan: &autoBan}).Error)
  76. }
  77. func putUserVideoChannelBindings(t *testing.T, router *gin.Engine, body string) *httptest.ResponseRecorder {
  78. t.Helper()
  79. w := httptest.NewRecorder()
  80. router.ServeHTTP(w, httptest.NewRequest(http.MethodPut, "/api/user/10/video-channel-bindings", bytes.NewBufferString(body)))
  81. return w
  82. }
  83. func TestAdminGetUserVideoChannelBindings(t *testing.T) {
  84. db := setupUserVideoChannelBindingDB(t)
  85. router := setupUserVideoChannelBindingRouter()
  86. require.NoError(t, db.Create(&model.User{Id: 10, Username: "user"}).Error)
  87. for i, group := range []string{"default", "vip", "auto"} {
  88. require.NoError(t, db.Create(&model.Token{Id: i + 1, UserId: 10, Key: "token-" + string(rune('0'+i)), Group: group}).Error)
  89. }
  90. createUserVideoBindingChannel(t, db, 7, constant.ChannelTypeDoubaoVideoCompatibleAiping, "default", common.ChannelStatusEnabled)
  91. createUserVideoBindingChannel(t, db, 16, constant.ChannelTypeDoubaoVideoCompatibleTianyiYun, "default", common.ChannelStatusEnabled)
  92. require.NoError(t, model.BindUserAssetChannel(10, constant.ChannelTypeDoubaoVideoCompatibleAiping, "default", 7))
  93. w := httptest.NewRecorder()
  94. router.ServeHTTP(w, httptest.NewRequest(http.MethodGet, "/api/user/10/video-channel-bindings", nil))
  95. require.Equal(t, http.StatusOK, w.Code)
  96. assert.Contains(t, w.Body.String(), `"default"`)
  97. assert.NotContains(t, w.Body.String(), `"vip"`)
  98. assert.NotContains(t, w.Body.String(), `"auto"`)
  99. assert.Contains(t, w.Body.String(), `"channel_id":7`)
  100. assert.NotContains(t, w.Body.String(), `channel-key`)
  101. }
  102. func TestAdminGetUserVideoChannelBindingsFiltersGroupsAndFamiliesWithoutCandidates(t *testing.T) {
  103. db := setupUserVideoChannelBindingDB(t)
  104. router := setupUserVideoChannelBindingRouter()
  105. require.NoError(t, db.Create(&model.User{Id: 10, Username: "user"}).Error)
  106. for i, group := range []string{"default", "vip", "test"} {
  107. require.NoError(t, db.Create(&model.Token{Id: i + 1, UserId: 10, Key: "token-" + group, Group: group}).Error)
  108. }
  109. createUserVideoBindingChannel(t, db, 7, constant.ChannelTypeDoubaoVideoCompatibleAiping, "default", common.ChannelStatusEnabled)
  110. createUserVideoBindingChannel(t, db, 59, constant.ChannelTypeKlingAiping, "default", common.ChannelStatusEnabled)
  111. createUserVideoBindingChannel(t, db, 16, constant.ChannelTypeDoubaoVideoCompatibleTianyiYun, "vip", common.ChannelStatusEnabled)
  112. w := httptest.NewRecorder()
  113. router.ServeHTTP(w, httptest.NewRequest(http.MethodGet, "/api/user/10/video-channel-bindings", nil))
  114. var response struct {
  115. Success bool `json:"success"`
  116. Data struct {
  117. Groups []string `json:"groups"`
  118. Families []adminVideoChannelBindingFamily `json:"families"`
  119. } `json:"data"`
  120. }
  121. require.NoError(t, common.Unmarshal(w.Body.Bytes(), &response))
  122. require.True(t, response.Success)
  123. assert.Equal(t, []string{"default", "vip"}, response.Data.Groups)
  124. familyGroups := map[string][]string{}
  125. for _, family := range response.Data.Families {
  126. for _, binding := range family.Bindings {
  127. familyGroups[family.Key] = append(familyGroups[family.Key], binding.Group)
  128. require.NotEmpty(t, binding.Candidates)
  129. }
  130. }
  131. assert.Equal(t, []string{"default", "vip"}, familyGroups["seedance"])
  132. assert.Equal(t, []string{"default"}, familyGroups["kling"])
  133. }
  134. func TestAdminSetUserVideoChannelBindings(t *testing.T) {
  135. db := setupUserVideoChannelBindingDB(t)
  136. router := setupUserVideoChannelBindingRouter()
  137. require.NoError(t, db.Create(&model.User{Id: 10, Username: "user"}).Error)
  138. require.NoError(t, db.Create(&model.Token{Id: 1, UserId: 10, Key: "token", Group: "default"}).Error)
  139. createUserVideoBindingChannel(t, db, 7, constant.ChannelTypeDoubaoVideoCompatibleAiping, "default", common.ChannelStatusEnabled)
  140. createUserVideoBindingChannel(t, db, 16, constant.ChannelTypeDoubaoVideoCompatibleTianyiYun, "default", common.ChannelStatusEnabled)
  141. createUserVideoBindingChannel(t, db, 59, constant.ChannelTypeKlingAiping, "default", common.ChannelStatusEnabled)
  142. require.NoError(t, model.BindUserAssetChannel(10, constant.ChannelTypeDoubaoVideoCompatibleAiping, "default", 7))
  143. require.NoError(t, model.BindUserAssetChannel(10, constant.ChannelTypeKlingAiping, "default", 59))
  144. w := putUserVideoChannelBindings(t, router, `{"bindings":[{"group":"default","family":"seedance","channel_id":16}]}`)
  145. require.Equal(t, http.StatusOK, w.Code)
  146. assert.Contains(t, w.Body.String(), `"success":true`)
  147. oldBinding, err := model.GetUserAssetChannel(10, constant.ChannelTypeDoubaoVideoCompatibleAiping, "default")
  148. require.NoError(t, err)
  149. assert.Nil(t, oldBinding)
  150. newBinding, err := model.GetUserAssetChannel(10, constant.ChannelTypeDoubaoVideoCompatibleTianyiYun, "default")
  151. require.NoError(t, err)
  152. require.NotNil(t, newBinding)
  153. assert.Equal(t, 16, newBinding.ChannelId)
  154. klingBinding, err := model.GetUserAssetChannel(10, constant.ChannelTypeKlingAiping, "default")
  155. require.NoError(t, err)
  156. assert.Nil(t, klingBinding)
  157. }
  158. func TestAdminSetUserVideoChannelBindingsRejectsInvalidRequestsWithoutChangingBindings(t *testing.T) {
  159. db := setupUserVideoChannelBindingDB(t)
  160. router := setupUserVideoChannelBindingRouter()
  161. require.NoError(t, db.Create(&model.User{Id: 10, Username: "user"}).Error)
  162. for i, group := range []string{"default", "auto"} {
  163. require.NoError(t, db.Create(&model.Token{Id: i + 1, UserId: 10, Key: "token-" + group, Group: group}).Error)
  164. }
  165. createUserVideoBindingChannel(t, db, 7, constant.ChannelTypeDoubaoVideoCompatibleAiping, "default", common.ChannelStatusEnabled)
  166. createUserVideoBindingChannel(t, db, 16, constant.ChannelTypeDoubaoVideoCompatibleTianyiYun, "default", common.ChannelStatusEnabled)
  167. createUserVideoBindingChannel(t, db, 59, constant.ChannelTypeKlingAiping, "default", common.ChannelStatusEnabled)
  168. createUserVideoBindingChannel(t, db, 60, constant.ChannelTypeDoubaoVideoCompatibleAiping, "vip", common.ChannelStatusEnabled)
  169. createUserVideoBindingChannel(t, db, 61, constant.ChannelTypeDoubaoVideoCompatibleAiping, "default", common.ChannelStatusAutoDisabled)
  170. require.NoError(t, model.BindUserAssetChannel(10, constant.ChannelTypeDoubaoVideoCompatibleAiping, "default", 7))
  171. for _, body := range []string{
  172. `{"bindings":[{"group":"vip","family":"seedance","channel_id":16}]}`,
  173. `{"bindings":[{"group":"auto","family":"seedance","channel_id":16}]}`,
  174. `{"bindings":[{"group":"default","family":"kling","channel_id":16}]}`,
  175. `{"bindings":[{"group":"default","family":"seedance","channel_id":60}]}`,
  176. `{"bindings":[{"group":"default","family":"seedance","channel_id":61}]}`,
  177. `{"bindings":[{"group":"default","family":"seedance","channel_id":16},{"group":"default","family":"seedance","channel_id":7}]}`,
  178. } {
  179. w := putUserVideoChannelBindings(t, router, body)
  180. assert.Contains(t, w.Body.String(), `"success":false`)
  181. binding, err := model.GetUserAssetChannel(10, constant.ChannelTypeDoubaoVideoCompatibleAiping, "default")
  182. require.NoError(t, err)
  183. require.NotNil(t, binding)
  184. assert.Equal(t, 7, binding.ChannelId)
  185. }
  186. }
  187. func TestAdminSetUserVideoChannelBindingsPreservesLegacyGroupBinding(t *testing.T) {
  188. db := setupUserVideoChannelBindingDB(t)
  189. router := setupUserVideoChannelBindingRouter()
  190. require.NoError(t, db.Create(&model.User{Id: 10, Username: "user"}).Error)
  191. require.NoError(t, db.Create(&model.Token{Id: 1, UserId: 10, Key: "token", Group: "default"}).Error)
  192. createUserVideoBindingChannel(t, db, 7, constant.ChannelTypeDoubaoVideoCompatibleAiping, "default", common.ChannelStatusEnabled)
  193. createUserVideoBindingChannel(t, db, 59, constant.ChannelTypeKlingAiping, "legacy", common.ChannelStatusEnabled)
  194. require.NoError(t, model.BindUserAssetChannel(10, constant.ChannelTypeKlingAiping, "legacy", 59))
  195. w := putUserVideoChannelBindings(t, router, `{"bindings":[{"group":"default","family":"seedance","channel_id":7}]}`)
  196. assert.Contains(t, w.Body.String(), `"success":true`)
  197. legacyBinding, err := model.GetUserAssetChannel(10, constant.ChannelTypeKlingAiping, "legacy")
  198. require.NoError(t, err)
  199. require.NotNil(t, legacyBinding)
  200. assert.Equal(t, 59, legacyBinding.ChannelId)
  201. }
  202. func TestAdminSetUserVideoChannelBindingsClearsHiddenCurrentGroupBinding(t *testing.T) {
  203. db := setupUserVideoChannelBindingDB(t)
  204. router := setupUserVideoChannelBindingRouter()
  205. require.NoError(t, db.Create(&model.User{Id: 10, Username: "user"}).Error)
  206. for i, group := range []string{"default", "test"} {
  207. require.NoError(t, db.Create(&model.Token{Id: i + 1, UserId: 10, Key: "token-" + group, Group: group}).Error)
  208. }
  209. createUserVideoBindingChannel(t, db, 7, constant.ChannelTypeDoubaoVideoCompatibleAiping, "default", common.ChannelStatusEnabled)
  210. require.NoError(t, model.BindUserAssetChannel(10, constant.ChannelTypeDoubaoVideoCompatibleAiping, "default", 7))
  211. require.NoError(t, model.BindUserAssetChannel(10, constant.ChannelTypeDoubaoVideoCompatibleAiping, "test", 99))
  212. w := putUserVideoChannelBindings(t, router, `{"bindings":[]}`)
  213. assert.Contains(t, w.Body.String(), `"success":true`)
  214. for _, group := range []string{"default", "test"} {
  215. binding, err := model.GetUserAssetChannel(10, constant.ChannelTypeDoubaoVideoCompatibleAiping, group)
  216. require.NoError(t, err)
  217. assert.Nil(t, binding)
  218. }
  219. }