package controller import ( "bytes" "net/http" "net/http/httptest" "os" "strings" "testing" "github.com/QuantumNous/new-api/common" "github.com/QuantumNous/new-api/constant" "github.com/QuantumNous/new-api/model" "github.com/gin-gonic/gin" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" "gorm.io/gorm" ) func setupUserVideoChannelBindingDB(t *testing.T) *gorm.DB { t.Helper() originalDB := model.DB originalLogDB := model.LOG_DB originalCache := common.MemoryCacheEnabled originalRedisEnabled := common.RedisEnabled originalSQLitePath := common.SQLitePath originalIsMasterNode := common.IsMasterNode originalUsingSQLite := common.UsingSQLite originalUsingMySQL := common.UsingMySQL originalUsingPostgreSQL := common.UsingPostgreSQL originalSQLDSN, hadSQLDSN := os.LookupEnv("SQL_DSN") common.SQLitePath = "file:" + strings.NewReplacer("/", "_", " ", "_").Replace(t.Name()) + "?mode=memory&cache=shared" common.MemoryCacheEnabled = false common.RedisEnabled = false common.IsMasterNode = false common.UsingSQLite = false common.UsingMySQL = false common.UsingPostgreSQL = false require.NoError(t, os.Setenv("SQL_DSN", "local")) require.NoError(t, model.InitDB()) db := model.DB model.LOG_DB = db sqlDB, err := db.DB() require.NoError(t, err) sqlDB.SetMaxOpenConns(1) require.NoError(t, db.AutoMigrate(&model.User{}, &model.Token{}, &model.Channel{}, &model.Ability{}, &model.UserAssetChannel{}, &model.Log{})) t.Cleanup(func() { model.DB = originalDB model.LOG_DB = originalLogDB common.MemoryCacheEnabled = originalCache common.RedisEnabled = originalRedisEnabled common.SQLitePath = originalSQLitePath common.IsMasterNode = originalIsMasterNode common.UsingSQLite = originalUsingSQLite common.UsingMySQL = originalUsingMySQL common.UsingPostgreSQL = originalUsingPostgreSQL if hadSQLDSN { _ = os.Setenv("SQL_DSN", originalSQLDSN) } else { _ = os.Unsetenv("SQL_DSN") } require.NoError(t, sqlDB.Close()) }) return db } func setupUserVideoChannelBindingRouter() *gin.Engine { gin.SetMode(gin.TestMode) router := gin.New() router.GET("/api/user/:id/video-channel-bindings", GetUserVideoChannelBindings) router.PUT("/api/user/:id/video-channel-bindings", SetUserVideoChannelBindings) return router } func createUserVideoBindingChannel(t *testing.T, db *gorm.DB, id, channelType int, group string, status int) { t.Helper() priority := int64(id) weight := uint(1) autoBan := 1 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) } func putUserVideoChannelBindings(t *testing.T, router *gin.Engine, body string) *httptest.ResponseRecorder { t.Helper() w := httptest.NewRecorder() router.ServeHTTP(w, httptest.NewRequest(http.MethodPut, "/api/user/10/video-channel-bindings", bytes.NewBufferString(body))) return w } func TestAdminGetUserVideoChannelBindings(t *testing.T) { db := setupUserVideoChannelBindingDB(t) router := setupUserVideoChannelBindingRouter() require.NoError(t, db.Create(&model.User{Id: 10, Username: "user"}).Error) for i, group := range []string{"default", "vip", "auto"} { require.NoError(t, db.Create(&model.Token{Id: i + 1, UserId: 10, Key: "token-" + string(rune('0'+i)), Group: group}).Error) } createUserVideoBindingChannel(t, db, 7, constant.ChannelTypeDoubaoVideoCompatibleAiping, "default", common.ChannelStatusEnabled) createUserVideoBindingChannel(t, db, 16, constant.ChannelTypeDoubaoVideoCompatibleTianyiYun, "default", common.ChannelStatusEnabled) require.NoError(t, model.BindUserAssetChannel(10, constant.ChannelTypeDoubaoVideoCompatibleAiping, "default", 7)) w := httptest.NewRecorder() router.ServeHTTP(w, httptest.NewRequest(http.MethodGet, "/api/user/10/video-channel-bindings", nil)) require.Equal(t, http.StatusOK, w.Code) assert.Contains(t, w.Body.String(), `"default"`) assert.NotContains(t, w.Body.String(), `"vip"`) assert.NotContains(t, w.Body.String(), `"auto"`) assert.Contains(t, w.Body.String(), `"channel_id":7`) assert.NotContains(t, w.Body.String(), `channel-key`) } func TestAdminGetUserVideoChannelBindingsFiltersGroupsAndFamiliesWithoutCandidates(t *testing.T) { db := setupUserVideoChannelBindingDB(t) router := setupUserVideoChannelBindingRouter() require.NoError(t, db.Create(&model.User{Id: 10, Username: "user"}).Error) for i, group := range []string{"default", "vip", "test"} { require.NoError(t, db.Create(&model.Token{Id: i + 1, UserId: 10, Key: "token-" + group, Group: group}).Error) } createUserVideoBindingChannel(t, db, 7, constant.ChannelTypeDoubaoVideoCompatibleAiping, "default", common.ChannelStatusEnabled) createUserVideoBindingChannel(t, db, 59, constant.ChannelTypeKlingAiping, "default", common.ChannelStatusEnabled) createUserVideoBindingChannel(t, db, 16, constant.ChannelTypeDoubaoVideoCompatibleTianyiYun, "vip", common.ChannelStatusEnabled) w := httptest.NewRecorder() router.ServeHTTP(w, httptest.NewRequest(http.MethodGet, "/api/user/10/video-channel-bindings", nil)) var response struct { Success bool `json:"success"` Data struct { Groups []string `json:"groups"` Families []adminVideoChannelBindingFamily `json:"families"` } `json:"data"` } require.NoError(t, common.Unmarshal(w.Body.Bytes(), &response)) require.True(t, response.Success) assert.Equal(t, []string{"default", "vip"}, response.Data.Groups) familyGroups := map[string][]string{} for _, family := range response.Data.Families { for _, binding := range family.Bindings { familyGroups[family.Key] = append(familyGroups[family.Key], binding.Group) require.NotEmpty(t, binding.Candidates) } } assert.Equal(t, []string{"default", "vip"}, familyGroups["seedance"]) assert.Equal(t, []string{"default"}, familyGroups["kling"]) } func TestAdminSetUserVideoChannelBindings(t *testing.T) { db := setupUserVideoChannelBindingDB(t) router := setupUserVideoChannelBindingRouter() require.NoError(t, db.Create(&model.User{Id: 10, Username: "user"}).Error) require.NoError(t, db.Create(&model.Token{Id: 1, UserId: 10, Key: "token", Group: "default"}).Error) createUserVideoBindingChannel(t, db, 7, constant.ChannelTypeDoubaoVideoCompatibleAiping, "default", common.ChannelStatusEnabled) createUserVideoBindingChannel(t, db, 16, constant.ChannelTypeDoubaoVideoCompatibleTianyiYun, "default", common.ChannelStatusEnabled) createUserVideoBindingChannel(t, db, 59, constant.ChannelTypeKlingAiping, "default", common.ChannelStatusEnabled) require.NoError(t, model.BindUserAssetChannel(10, constant.ChannelTypeDoubaoVideoCompatibleAiping, "default", 7)) require.NoError(t, model.BindUserAssetChannel(10, constant.ChannelTypeKlingAiping, "default", 59)) w := putUserVideoChannelBindings(t, router, `{"bindings":[{"group":"default","family":"seedance","channel_id":16}]}`) require.Equal(t, http.StatusOK, w.Code) assert.Contains(t, w.Body.String(), `"success":true`) oldBinding, err := model.GetUserAssetChannel(10, constant.ChannelTypeDoubaoVideoCompatibleAiping, "default") require.NoError(t, err) assert.Nil(t, oldBinding) newBinding, err := model.GetUserAssetChannel(10, constant.ChannelTypeDoubaoVideoCompatibleTianyiYun, "default") require.NoError(t, err) require.NotNil(t, newBinding) assert.Equal(t, 16, newBinding.ChannelId) klingBinding, err := model.GetUserAssetChannel(10, constant.ChannelTypeKlingAiping, "default") require.NoError(t, err) assert.Nil(t, klingBinding) } func TestAdminSetUserVideoChannelBindingsRejectsInvalidRequestsWithoutChangingBindings(t *testing.T) { db := setupUserVideoChannelBindingDB(t) router := setupUserVideoChannelBindingRouter() require.NoError(t, db.Create(&model.User{Id: 10, Username: "user"}).Error) for i, group := range []string{"default", "auto"} { require.NoError(t, db.Create(&model.Token{Id: i + 1, UserId: 10, Key: "token-" + group, Group: group}).Error) } createUserVideoBindingChannel(t, db, 7, constant.ChannelTypeDoubaoVideoCompatibleAiping, "default", common.ChannelStatusEnabled) createUserVideoBindingChannel(t, db, 16, constant.ChannelTypeDoubaoVideoCompatibleTianyiYun, "default", common.ChannelStatusEnabled) createUserVideoBindingChannel(t, db, 59, constant.ChannelTypeKlingAiping, "default", common.ChannelStatusEnabled) createUserVideoBindingChannel(t, db, 60, constant.ChannelTypeDoubaoVideoCompatibleAiping, "vip", common.ChannelStatusEnabled) createUserVideoBindingChannel(t, db, 61, constant.ChannelTypeDoubaoVideoCompatibleAiping, "default", common.ChannelStatusAutoDisabled) require.NoError(t, model.BindUserAssetChannel(10, constant.ChannelTypeDoubaoVideoCompatibleAiping, "default", 7)) for _, body := range []string{ `{"bindings":[{"group":"vip","family":"seedance","channel_id":16}]}`, `{"bindings":[{"group":"auto","family":"seedance","channel_id":16}]}`, `{"bindings":[{"group":"default","family":"kling","channel_id":16}]}`, `{"bindings":[{"group":"default","family":"seedance","channel_id":60}]}`, `{"bindings":[{"group":"default","family":"seedance","channel_id":61}]}`, `{"bindings":[{"group":"default","family":"seedance","channel_id":16},{"group":"default","family":"seedance","channel_id":7}]}`, } { w := putUserVideoChannelBindings(t, router, body) assert.Contains(t, w.Body.String(), `"success":false`) binding, err := model.GetUserAssetChannel(10, constant.ChannelTypeDoubaoVideoCompatibleAiping, "default") require.NoError(t, err) require.NotNil(t, binding) assert.Equal(t, 7, binding.ChannelId) } } func TestAdminSetUserVideoChannelBindingsPreservesLegacyGroupBinding(t *testing.T) { db := setupUserVideoChannelBindingDB(t) router := setupUserVideoChannelBindingRouter() require.NoError(t, db.Create(&model.User{Id: 10, Username: "user"}).Error) require.NoError(t, db.Create(&model.Token{Id: 1, UserId: 10, Key: "token", Group: "default"}).Error) createUserVideoBindingChannel(t, db, 7, constant.ChannelTypeDoubaoVideoCompatibleAiping, "default", common.ChannelStatusEnabled) createUserVideoBindingChannel(t, db, 59, constant.ChannelTypeKlingAiping, "legacy", common.ChannelStatusEnabled) require.NoError(t, model.BindUserAssetChannel(10, constant.ChannelTypeKlingAiping, "legacy", 59)) w := putUserVideoChannelBindings(t, router, `{"bindings":[{"group":"default","family":"seedance","channel_id":7}]}`) assert.Contains(t, w.Body.String(), `"success":true`) legacyBinding, err := model.GetUserAssetChannel(10, constant.ChannelTypeKlingAiping, "legacy") require.NoError(t, err) require.NotNil(t, legacyBinding) assert.Equal(t, 59, legacyBinding.ChannelId) } func TestAdminSetUserVideoChannelBindingsClearsHiddenCurrentGroupBinding(t *testing.T) { db := setupUserVideoChannelBindingDB(t) router := setupUserVideoChannelBindingRouter() require.NoError(t, db.Create(&model.User{Id: 10, Username: "user"}).Error) for i, group := range []string{"default", "test"} { require.NoError(t, db.Create(&model.Token{Id: i + 1, UserId: 10, Key: "token-" + group, Group: group}).Error) } createUserVideoBindingChannel(t, db, 7, constant.ChannelTypeDoubaoVideoCompatibleAiping, "default", common.ChannelStatusEnabled) require.NoError(t, model.BindUserAssetChannel(10, constant.ChannelTypeDoubaoVideoCompatibleAiping, "default", 7)) require.NoError(t, model.BindUserAssetChannel(10, constant.ChannelTypeDoubaoVideoCompatibleAiping, "test", 99)) w := putUserVideoChannelBindings(t, router, `{"bindings":[]}`) assert.Contains(t, w.Body.String(), `"success":true`) for _, group := range []string{"default", "test"} { binding, err := model.GetUserAssetChannel(10, constant.ChannelTypeDoubaoVideoCompatibleAiping, group) require.NoError(t, err) assert.Nil(t, binding) } }