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