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.

112 lines
5.5 KiB

  1. package service
  2. import (
  3. "testing"
  4. "github.com/QuantumNous/new-api/common"
  5. "github.com/QuantumNous/new-api/constant"
  6. "github.com/QuantumNous/new-api/model"
  7. "github.com/stretchr/testify/assert"
  8. "github.com/stretchr/testify/require"
  9. )
  10. func TestGetBoundVideoAssetChannelForModelFindsSeedanceTianyiYunBinding(t *testing.T) {
  11. db := setupDoubaoAssetChannelDB(t)
  12. createDoubaoAssetChannelForTest(t, db, 16, constant.ChannelTypeDoubaoVideoCompatibleTianyiYun, "default", "key", common.ChannelStatusEnabled)
  13. createDoubaoAssetAbilityForTest(t, db, "default", "Doubao-Seedance-2.0", 16, true)
  14. require.NoError(t, model.BindUserAssetChannel(10, constant.ChannelTypeDoubaoVideoCompatibleTianyiYun, "default", 16))
  15. ch, err := GetBoundVideoAssetChannelForModel(10, "default", "Doubao-Seedance-2.0", VideoAssetFamilySeedance)
  16. require.NoError(t, err)
  17. require.NotNil(t, ch)
  18. assert.Equal(t, 16, ch.Id)
  19. }
  20. func TestVideoAssetSeedanceFamilyIncludesChinaMobileSeedance(t *testing.T) {
  21. require.Contains(t, VideoAssetChannelTypesForFamily(VideoAssetFamilySeedance), constant.ChannelTypeChinaMobileSeedance)
  22. }
  23. func TestGetBoundVideoAssetChannelForModelFindsKlingBindingByFamily(t *testing.T) {
  24. db := setupDoubaoAssetChannelDB(t)
  25. createDoubaoAssetChannelForTest(t, db, 59, constant.ChannelTypeKlingAiping, "default", "key", common.ChannelStatusEnabled)
  26. createDoubaoAssetAbilityForTest(t, db, "default", "kling-v2-6", 59, true)
  27. require.NoError(t, model.BindUserAssetChannel(10, constant.ChannelTypeKlingAiping, "default", 59))
  28. ch, err := GetBoundVideoAssetChannelForModel(10, "default", "kling-v2-6", VideoAssetFamilyKling)
  29. require.NoError(t, err)
  30. require.NotNil(t, ch)
  31. assert.Equal(t, 59, ch.Id)
  32. }
  33. func TestGetBoundVideoAssetChannelForModelRejectsBindingWithoutAbility(t *testing.T) {
  34. db := setupDoubaoAssetChannelDB(t)
  35. createDoubaoAssetChannelForTest(t, db, 16, constant.ChannelTypeDoubaoVideoCompatibleTianyiYun, "default", "key", common.ChannelStatusEnabled)
  36. require.NoError(t, model.BindUserAssetChannel(10, constant.ChannelTypeDoubaoVideoCompatibleTianyiYun, "default", 16))
  37. ch, err := GetBoundVideoAssetChannelForModel(10, "default", "Doubao-Seedance-2.0", VideoAssetFamilySeedance)
  38. require.NoError(t, err)
  39. assert.Nil(t, ch)
  40. }
  41. func TestBindVideoAssetChannelReplacesOtherFamilyBindings(t *testing.T) {
  42. db := setupDoubaoAssetChannelDB(t)
  43. createDoubaoAssetChannelForTest(t, db, 7, constant.ChannelTypeDoubaoVideoCompatibleAiping, "default", "aiping-key", common.ChannelStatusEnabled)
  44. createDoubaoAssetChannelForTest(t, db, 16, constant.ChannelTypeDoubaoVideoCompatibleTianyiYun, "default", "tianyiyun-key", common.ChannelStatusEnabled)
  45. createDoubaoAssetAbilityForTest(t, db, "default", "Doubao-Seedance-2.0", 7, true)
  46. createDoubaoAssetAbilityForTest(t, db, "default", "Doubao-Seedance-2.0", 16, true)
  47. require.NoError(t, model.BindUserAssetChannel(10, constant.ChannelTypeDoubaoVideoCompatibleAiping, "default", 7))
  48. selected, err := model.CacheGetChannel(16)
  49. require.NoError(t, err)
  50. require.NoError(t, BindVideoAssetChannel(10, "default", selected, VideoAssetFamilySeedance))
  51. oldBinding, err := model.GetUserAssetChannel(10, constant.ChannelTypeDoubaoVideoCompatibleAiping, "default")
  52. require.NoError(t, err)
  53. assert.Nil(t, oldBinding)
  54. newBinding, err := model.GetUserAssetChannel(10, constant.ChannelTypeDoubaoVideoCompatibleTianyiYun, "default")
  55. require.NoError(t, err)
  56. require.NotNil(t, newBinding)
  57. assert.Equal(t, 16, newBinding.ChannelId)
  58. }
  59. func TestGetBoundVideoAssetChannelForModelPrefersLatestFamilyBindingForLegacyRows(t *testing.T) {
  60. db := setupDoubaoAssetChannelDB(t)
  61. createDoubaoAssetChannelForTest(t, db, 7, constant.ChannelTypeDoubaoVideoCompatibleAiping, "default", "aiping-key", common.ChannelStatusEnabled)
  62. createDoubaoAssetChannelForTest(t, db, 16, constant.ChannelTypeDoubaoVideoCompatibleTianyiYun, "default", "tianyiyun-key", common.ChannelStatusEnabled)
  63. createDoubaoAssetAbilityForTest(t, db, "default", "Doubao-Seedance-2.0", 7, true)
  64. createDoubaoAssetAbilityForTest(t, db, "default", "Doubao-Seedance-2.0", 16, true)
  65. require.NoError(t, model.BindUserAssetChannel(10, constant.ChannelTypeDoubaoVideoCompatibleAiping, "default", 7))
  66. require.NoError(t, model.BindUserAssetChannel(10, constant.ChannelTypeDoubaoVideoCompatibleTianyiYun, "default", 16))
  67. require.NoError(t, db.Model(&model.UserAssetChannel{}).
  68. Where("user_id = ? AND channel_type = ?", 10, constant.ChannelTypeDoubaoVideoCompatibleAiping).
  69. Update("updated_at", int64(100)).Error)
  70. require.NoError(t, db.Model(&model.UserAssetChannel{}).
  71. Where("user_id = ? AND channel_type = ?", 10, constant.ChannelTypeDoubaoVideoCompatibleTianyiYun).
  72. Update("updated_at", int64(200)).Error)
  73. ch, err := GetBoundVideoAssetChannelForModel(10, "default", "Doubao-Seedance-2.0", VideoAssetFamilySeedance)
  74. require.NoError(t, err)
  75. require.NotNil(t, ch)
  76. assert.Equal(t, 16, ch.Id)
  77. }
  78. func TestResolveVideoAssetChannelForModelAutoSelectsAndPersistsFamilyBinding(t *testing.T) {
  79. db := setupDoubaoAssetChannelDB(t)
  80. createDoubaoAssetChannelForTest(t, db, 16, constant.ChannelTypeDoubaoVideoCompatibleTianyiYun, "default", "tianyiyun-key", common.ChannelStatusEnabled)
  81. createDoubaoAssetAbilityForTest(t, db, "default", "Doubao-Seedance-2.0", 16, true)
  82. ch, err := ResolveVideoAssetChannelForModel(10, "default", "Doubao-Seedance-2.0", VideoAssetFamilySeedance)
  83. require.NoError(t, err)
  84. require.NotNil(t, ch)
  85. assert.Equal(t, 16, ch.Id)
  86. binding, err := model.GetUserAssetChannel(10, constant.ChannelTypeDoubaoVideoCompatibleTianyiYun, "default")
  87. require.NoError(t, err)
  88. require.NotNil(t, binding)
  89. assert.Equal(t, 16, binding.ChannelId)
  90. }