package router import ( "net/http" "net/http/httptest" "os" "strings" "sync/atomic" "testing" "github.com/QuantumNous/new-api/common" "github.com/QuantumNous/new-api/constant" "github.com/QuantumNous/new-api/model" "github.com/QuantumNous/new-api/service" "github.com/QuantumNous/new-api/setting/ratio_setting" "github.com/gin-gonic/gin" "github.com/stretchr/testify/require" "gorm.io/gorm" "gorm.io/gorm/logger" ) func setupTianyiYunSeedanceE2EDB(t *testing.T) *gorm.DB { t.Helper() oldDB := model.DB oldLOGDB := model.LOG_DB oldSQLitePath := common.SQLitePath oldMemoryCacheEnabled := common.MemoryCacheEnabled oldRedisEnabled := common.RedisEnabled oldIsMasterNode := common.IsMasterNode oldUsingSQLite := common.UsingSQLite oldUsingMySQL := common.UsingMySQL oldUsingPostgreSQL := common.UsingPostgreSQL oldSQLDSN, hadSQLDSN := os.LookupEnv("SQL_DSN") oldModelRatio := ratio_setting.ModelRatio2JSONString() oldGroupRatio := ratio_setting.GroupRatio2JSONString() 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()) require.NoError(t, ratio_setting.UpdateModelRatioByJSONString(`{"cdance2.0-0611":0}`)) require.NoError(t, ratio_setting.UpdateGroupRatioByJSONString(`{"default":1}`)) service.InitHttpClient() model.DB = model.DB.Session(&gorm.Session{Logger: logger.Default.LogMode(logger.Silent)}) model.LOG_DB = model.DB db := model.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.Task{}, &model.Log{}, &model.SubscriptionPlan{}, &model.SubscriptionOrder{}, &model.UserSubscription{}, &model.SubscriptionPreConsumeRecord{}, )) t.Cleanup(func() { _ = sqlDB.Close() model.DB = oldDB model.LOG_DB = oldLOGDB common.SQLitePath = oldSQLitePath common.MemoryCacheEnabled = oldMemoryCacheEnabled common.RedisEnabled = oldRedisEnabled common.IsMasterNode = oldIsMasterNode common.UsingSQLite = oldUsingSQLite common.UsingMySQL = oldUsingMySQL common.UsingPostgreSQL = oldUsingPostgreSQL _ = ratio_setting.UpdateModelRatioByJSONString(oldModelRatio) _ = ratio_setting.UpdateGroupRatioByJSONString(oldGroupRatio) if hadSQLDSN { _ = os.Setenv("SQL_DSN", oldSQLDSN) } else { _ = os.Unsetenv("SQL_DSN") } }) return db } func TestTianyiYunSeedanceSubmitE2EUsesBoundUserChannel(t *testing.T) { gin.SetMode(gin.TestMode) db := setupTianyiYunSeedanceE2EDB(t) var upstreamCalls int32 var gotPath string var gotAuth string upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { atomic.AddInt32(&upstreamCalls, 1) gotPath = r.URL.Path gotAuth = r.Header.Get("Authorization") w.Header().Set("Content-Type", "application/json") _, _ = w.Write([]byte(`{"id":"task-upstream","status":"queued","model":"cdance2.0-0611"}`)) })) t.Cleanup(upstream.Close) require.NoError(t, db.Create(&model.User{ Id: 10, Username: "e2e-user", Status: common.UserStatusEnabled, Group: "default", Quota: 1000000, }).Error) require.NoError(t, db.Create(&model.Token{ Id: 20, UserId: 10, Key: "e2etokenkey", Status: common.TokenStatusEnabled, Name: "e2e-token", ExpiredTime: -1, UnlimitedQuota: true, Group: "default", }).Error) priority := int64(1) weight := uint(10) autoBan := 1 baseURL := upstream.URL require.NoError(t, db.Create(&model.Channel{ Id: 30, Type: constant.ChannelTypeDoubaoVideoCompatibleTianyiYun, Key: "upstream-secret", Status: common.ChannelStatusEnabled, Name: "tianyiyun-e2e", Group: "default", Models: "cdance2.0-0611", BaseURL: &baseURL, Priority: &priority, Weight: &weight, AutoBan: &autoBan, CreatedTime: 30, }).Error) require.NoError(t, db.Create(&model.Ability{ Group: "default", Model: "cdance2.0-0611", ChannelId: 30, Enabled: true, Priority: &priority, Weight: 10, }).Error) require.NoError(t, model.BindUserAssetChannel(10, constant.ChannelTypeDoubaoVideoCompatibleTianyiYun, "default", 30)) r := gin.New() SetVideoRouter(r) body := `{ "model":"cdance2.0-0611", "content":[{"type":"text","text":"e2e prompt"}], "ratio":"16:9", "duration":5, "watermark":false }` req := httptest.NewRequest(http.MethodPost, "/api/v3/contents/generations/tasks", strings.NewReader(body)) req.Header.Set("Content-Type", "application/json") req.Header.Set("Authorization", "Bearer sk-e2etokenkey") w := httptest.NewRecorder() r.ServeHTTP(w, req) require.Equal(t, http.StatusOK, w.Code, w.Body.String()) require.Equal(t, int32(1), atomic.LoadInt32(&upstreamCalls)) require.Equal(t, "/v1/contents/generations/tasks", gotPath) require.Equal(t, "Bearer upstream-secret", gotAuth) require.Contains(t, w.Body.String(), `"id":"task_`) binding, err := model.GetUserAssetChannel(10, constant.ChannelTypeDoubaoVideoCompatibleTianyiYun, "default") require.NoError(t, err) require.NotNil(t, binding) require.Equal(t, 30, binding.ChannelId) var task model.Task require.NoError(t, db.Where("user_id = ? AND channel_id = ?", 10, 30).First(&task).Error) require.Equal(t, constant.TaskPlatform("60"), task.Platform) require.Equal(t, "task-upstream", task.PrivateData.UpstreamTaskID) require.Equal(t, "cdance2.0-0611", task.Properties.OriginModelName) }