11 коммитов

Автор SHA1 Сообщение Дата
  fengsilin 83724d4930 chore: add image build script, drop binary artifact, add missing go.sum 1 месяц назад
  fengsilin ce40bfe92f feat(chinamobile): support doubao-seedance-2-0-mini-260615 billing model 1 месяц назад
  fengsilin 0b79e4ee07 fix(db): skip logs AutoMigrate for pg_partman-managed PostgreSQL tables 1 месяц назад
  fengsilin 17c12fd41e feat: harden services and add video channel bindings 2 месяцев назад
  fengsilin cd35e66b33 feat(chinamobile): add channel asset credentials 2 месяцев назад
  fengsilin 50ce4f7541 feat(tianyiyun): add Seedance asset support 2 месяцев назад
  fengsilin fedd394e79 fix(task-billing): persist settlement and merge paged logs 2 месяцев назад
  fengsilin c67d8c409e feat(chinamobile): add Seedance channel and isolated asset library 2 месяцев назад
  fengsilin fcb1c2a8ba fix: prevent webhook bypass via empty secret and cross-gateway callback attacks 2 месяцев назад
  fengsilin 6891b11e09 chore: ignore local tmp artifacts 2 месяцев назад
  fengsilin bde79b7259 fix(kling): persist proxy fallback asset binding 2 месяцев назад
100 измененных файлов: 7300 добавлений и 194 удалений
  1. +4
    -0
      .gitignore
  2. +3
    -0
      Dockerfile
  3. +2
    -0
      common/api_type.go
  4. +17
    -0
      common/chinamobile_seedance_channel_test.go
  5. +1
    -1
      common/endpoint_type.go
  6. +3
    -0
      constant/channel.go
  7. +1
    -0
      controller/channel-test.go
  8. +101
    -9
      controller/channel.go
  9. +35
    -0
      controller/channel_affinity_cache_test.go
  10. +50
    -0
      controller/channel_asset_credential_test.go
  11. +25
    -0
      controller/channel_billing_helpers_test.go
  12. +45
    -0
      controller/deployment_test.go
  13. +74
    -120
      controller/doubao_asset.go
  14. +189
    -0
      controller/doubao_asset_real_e2e_test.go
  15. +330
    -14
      controller/doubao_asset_test.go
  16. +18
    -0
      controller/helpers_test.go
  17. +16
    -0
      controller/kling_aiping_native.go
  18. +113
    -0
      controller/kling_aiping_native_test.go
  19. +40
    -0
      controller/option_test.go
  20. +55
    -0
      controller/payment_webhook_availability.go
  21. +57
    -0
      controller/payment_webhook_availability_test.go
  22. +31
    -0
      controller/ratio_config_test.go
  23. +5
    -0
      controller/subscription_payment_epay.go
  24. +8
    -1
      controller/topup.go
  25. +8
    -1
      controller/topup_alipay.go
  26. +13
    -6
      controller/topup_creem.go
  27. +8
    -1
      controller/topup_stripe.go
  28. +8
    -1
      controller/topup_wechat.go
  29. +10
    -0
      controller/user.go
  30. +209
    -0
      controller/user_video_channel_binding.go
  31. +243
    -0
      controller/user_video_channel_binding_test.go
  32. +40
    -0
      controller/video_proxy_gemini_test.go
  33. +14
    -0
      deploy/build_newapi.sh
  34. +39
    -0
      deploy/postgres-partman/Dockerfile
  35. +102
    -0
      deploy/postgres-partman/docker-entrypoint-initdb.d/01-partition.sql
  36. +89
    -0
      docs/superpowers/specs/2026-07-17-chinamobile-asset-isolation-design.md
  37. +329
    -0
      docs/testing/2026-07-21-cn-tianyiyun-seedance-mini-e2e.md
  38. +30
    -0
      docs/testing/2026-07-24-admin-video-channel-binding.md
  39. +13
    -1
      go.mod
  40. +80
    -1
      go.sum
  41. +64
    -0
      middleware/gzip_test.go
  42. +44
    -0
      middleware/http_headers_test.go
  43. +41
    -0
      middleware/recover_test.go
  44. +86
    -16
      model/channel.go
  45. +99
    -0
      model/channel_asset_credential.go
  46. +132
    -0
      model/channel_asset_credential_test.go
  47. +44
    -0
      model/custom_oauth_provider_test.go
  48. +73
    -0
      model/log.go
  49. +50
    -0
      model/log_identity_test.go
  50. +73
    -3
      model/main.go
  51. +45
    -0
      model/option_map_test.go
  52. +42
    -0
      model/passkey_test.go
  53. +41
    -0
      model/prefill_group_test.go
  54. +13
    -0
      model/pricing_default_test.go
  55. +20
    -1
      model/token.go
  56. +38
    -0
      model/token_helpers_test.go
  57. +28
    -11
      model/topup.go
  58. +5
    -1
      model/topup_wechat.go
  59. +21
    -0
      model/twofa_test.go
  60. +8
    -0
      model/user_asset_channel.go
  61. +20
    -0
      model/user_asset_channel_test.go
  62. +45
    -0
      model/user_asset_group.go
  63. +63
    -0
      model/user_asset_group_test.go
  64. +42
    -0
      model/user_cache_test.go
  65. +38
    -0
      model/utils_test.go
  66. Двоичные данные
      new-api.exe~
  67. +494
    -0
      relay/channel/task/chinamobile_seedance/adaptor.go
  68. +349
    -0
      relay/channel/task/chinamobile_seedance/adaptor_test.go
  69. +9
    -0
      relay/channel/task/chinamobile_seedance/constants.go
  70. +81
    -0
      relay/channel/task/chinamobile_seedance/sdk_client.go
  71. +24
    -0
      relay/channel/task/doubao_tianyiyun/adaptor.go
  72. +51
    -0
      relay/channel/task/doubao_tianyiyun/adaptor_test.go
  73. +1
    -0
      relay/channel/task/doubao_tianyiyun/constants.go
  74. +9
    -0
      relay/helper/matrix_usage_capability.go
  75. +18
    -0
      relay/helper/matrix_usage_capability_test.go
  76. +3
    -0
      relay/relay_adaptor.go
  77. +15
    -0
      relay/relay_adaptor_chinamobile_test.go
  78. +5
    -0
      relay/relay_task.go
  79. +2
    -0
      router/api-router.go
  80. +101
    -0
      service/asset_adapter.go
  81. +28
    -0
      service/asset_adapter_test.go
  82. +758
    -0
      service/asset_chinamobile.go
  83. +41
    -0
      service/asset_chinamobile_sdk.go
  84. +19
    -0
      service/asset_chinamobile_sdk_test.go
  85. +467
    -0
      service/asset_chinamobile_test.go
  86. +98
    -0
      service/asset_compatible.go
  87. +27
    -0
      service/asset_compatible_test.go
  88. +77
    -0
      service/asset_operation.go
  89. +45
    -0
      service/asset_operation_test.go
  90. +79
    -0
      service/asset_resolver.go
  91. +106
    -0
      service/asset_resolver_test.go
  92. +163
    -0
      service/asset_tianyiyun.go
  93. +131
    -0
      service/asset_tianyiyun_test.go
  94. +51
    -0
      service/audio_test.go
  95. +35
    -0
      service/channel_affinity_helpers_test.go
  96. +20
    -0
      service/channel_select_test.go
  97. +67
    -0
      service/channel_test.go
  98. +128
    -0
      service/chinamobile_user_asset_group.go
  99. +159
    -0
      service/chinamobile_user_asset_group_test.go
  100. +6
    -6
      service/codex_oauth.go

+ 4
- 0
.gitignore Просмотреть файл

@@ -27,12 +27,16 @@ CLAUDE.md
.worktrees/
logs/
docs/superpowers
/tmp/

# Runtime request/response capture files
/[0-9]*.json

# Local dev server logs
web/vite-*.log
/server_*.log
/server_*_err.log
/asset_test_log.md

# E2E / Playwright artifacts
test-artifacts/


+ 3
- 0
Dockerfile Просмотреть файл

@@ -23,6 +23,9 @@ ENV GOOS=${TARGETOS:-linux} GOARCH=${TARGETARCH:-amd64}
WORKDIR /build

ADD go.mod go.sum ./
COPY ./third_party/ecloudsdkcore ./third_party/ecloudsdkcore
COPY ./third_party/ecloudsdkmaas ./third_party/ecloudsdkmaas
COPY ./third_party/maas_seedance_sdk_1.0.0_go ./third_party/maas_seedance_sdk_1.0.0_go
RUN --mount=type=cache,target=/go/pkg/mod go mod download

COPY . .


+ 2
- 0
common/api_type.go Просмотреть файл

@@ -57,6 +57,8 @@ func ChannelType2APIType(channelType int) (int, bool) {
apiType = constant.APITypeVolcEngine
case constant.ChannelTypeDoubaoVideoCompatibleTianyiYun:
apiType = constant.APITypeVolcEngine
case constant.ChannelTypeChinaMobileSeedance:
apiType = constant.APITypeVolcEngine
case constant.ChannelTypeBaiduV2:
apiType = constant.APITypeBaiduV2
case constant.ChannelTypeOpenRouter:


+ 17
- 0
common/chinamobile_seedance_channel_test.go Просмотреть файл

@@ -0,0 +1,17 @@
package common

import (
"testing"

"github.com/QuantumNous/new-api/constant"
"github.com/stretchr/testify/require"
)

func TestChinaMobileSeedanceChannelMappings(t *testing.T) {
apiType, ok := ChannelType2APIType(constant.ChannelTypeChinaMobileSeedance)
require.True(t, ok)
require.Equal(t, constant.APITypeVolcEngine, apiType)

endpoints := GetEndpointTypesByChannelType(constant.ChannelTypeChinaMobileSeedance, "doubao-seedance-2.0")
require.Contains(t, endpoints, constant.EndpointTypeDoubaoVideo)
}

+ 1
- 1
common/endpoint_type.go Просмотреть файл

@@ -30,7 +30,7 @@ func GetEndpointTypesByChannelType(channelType int, modelName string) []constant
endpointTypes = []constant.EndpointType{constant.EndpointTypeOpenAI, constant.EndpointTypeOpenAIResponse}
case constant.ChannelTypeSora:
endpointTypes = []constant.EndpointType{constant.EndpointTypeOpenAIVideo}
case constant.ChannelTypeDoubaoVideo, constant.ChannelTypeDoubaoVideoCompatibleAiping, constant.ChannelTypeDoubaoVideoCompatibleTianyiYun:
case constant.ChannelTypeDoubaoVideo, constant.ChannelTypeDoubaoVideoCompatibleAiping, constant.ChannelTypeDoubaoVideoCompatibleTianyiYun, constant.ChannelTypeChinaMobileSeedance:
endpointTypes = []constant.EndpointType{constant.EndpointTypeDoubaoVideo}
default:
if IsOpenAIResponseOnlyModel(modelName) {


+ 3
- 0
constant/channel.go Просмотреть файл

@@ -58,6 +58,7 @@ const (
ChannelTypeDoubaoVideoCompatibleAiping = 58
ChannelTypeKlingAiping = 59
ChannelTypeDoubaoVideoCompatibleTianyiYun = 60
ChannelTypeChinaMobileSeedance = 61
ChannelTypeDummy // this one is only for count, do not add any channel after this

)
@@ -124,6 +125,7 @@ var ChannelBaseURLs = []string{
"", //58
"https://aiping.cn/api", //59
"https://ai.ctaigw.cn", //60
"https://zhenze-huhehaote.cmecloud.cn/api/v3", //61
}

var ChannelTypeNames = map[int]string{
@@ -184,6 +186,7 @@ var ChannelTypeNames = map[int]string{
ChannelTypeDoubaoVideoCompatibleAiping: "DoubaoVideoCompatibleAiping",
ChannelTypeKlingAiping: "KlingAiping",
ChannelTypeDoubaoVideoCompatibleTianyiYun: "DoubaoVideoCompatibleTianyiYun",
ChannelTypeChinaMobileSeedance: "ChinaMobileSeedance",
}

func GetChannelTypeName(channelType int) string {


+ 1
- 0
controller/channel-test.go Просмотреть файл

@@ -67,6 +67,7 @@ func testChannel(channel *model.Channel, testModel string, endpointType string,
constant.ChannelTypeDoubaoVideo,
constant.ChannelTypeDoubaoVideoCompatibleAiping,
constant.ChannelTypeDoubaoVideoCompatibleTianyiYun,
constant.ChannelTypeChinaMobileSeedance,
constant.ChannelTypeVidu,
}
if lo.Contains(unsupportedTestChannelTypes, channel.Type) {


+ 101
- 9
controller/channel.go Просмотреть файл

@@ -19,6 +19,7 @@ import (
"github.com/QuantumNous/new-api/service"

"github.com/gin-gonic/gin"
"gorm.io/gorm"
)

type OpenAIModel struct {
@@ -68,6 +69,30 @@ func clearChannelInfo(channel *model.Channel) {
}
}

func attachChannelAssetCredentialSummaries(channels []*model.Channel) error {
ids := make([]int, 0)
for _, channel := range channels {
if channel != nil && channel.Type == constant.ChannelTypeChinaMobileSeedance {
ids = append(ids, channel.Id)
}
}
summaries, err := model.GetChannelAssetCredentialSummaries(ids)
if err != nil {
return err
}
for _, channel := range channels {
if channel == nil || channel.Type != constant.ChannelTypeChinaMobileSeedance {
continue
}
summary, ok := summaries[channel.Id]
channel.AssetCredentialConfigured = ok
if ok {
channel.AssetCredentialPoolID = summary.PoolID
}
}
return nil
}

func GetAllChannels(c *gin.Context) {
pageInfo := common.GetPageQuery(c)
channelData := make([]*model.Channel, 0)
@@ -144,6 +169,10 @@ func GetAllChannels(c *gin.Context) {
}
}

if err := attachChannelAssetCredentialSummaries(channelData); err != nil {
common.ApiError(c, err)
return
}
for _, datum := range channelData {
clearChannelInfo(datum)
}
@@ -482,6 +511,10 @@ func SearchChannels(c *gin.Context) {

pagedData := channelData[startIdx:endIdx]

if err := attachChannelAssetCredentialSummaries(pagedData); err != nil {
common.ApiError(c, err)
return
}
for _, datum := range pagedData {
clearChannelInfo(datum)
}
@@ -510,6 +543,10 @@ func GetChannel(c *gin.Context) {
return
}
if channel != nil {
if err := attachChannelAssetCredentialSummaries([]*model.Channel{channel}); err != nil {
common.ApiError(c, err)
return
}
clearChannelInfo(channel)
}
c.JSON(http.StatusOK, gin.H{
@@ -675,10 +712,29 @@ func RefreshCodexChannelCredential(c *gin.Context) {
}

type AddChannelRequest struct {
Mode string `json:"mode"`
MultiKeyMode constant.MultiKeyMode `json:"multi_key_mode"`
BatchAddSetKeyPrefix2Name bool `json:"batch_add_set_key_prefix_2_name"`
Channel *model.Channel `json:"channel"`
Mode string `json:"mode"`
MultiKeyMode constant.MultiKeyMode `json:"multi_key_mode"`
BatchAddSetKeyPrefix2Name bool `json:"batch_add_set_key_prefix_2_name"`
Channel *model.Channel `json:"channel"`
AssetCredential *ChannelAssetCredentialInput `json:"asset_credential"`
}

type ChannelAssetCredentialInput struct {
AccessKey string `json:"access_key"`
SecretKey string `json:"secret_key"`
PoolID string `json:"pool_id"`
}

func channelAssetCredentialFromInput(channelType int, input *ChannelAssetCredentialInput) (*model.ChannelAssetCredential, error) {
if channelType != constant.ChannelTypeChinaMobileSeedance || input == nil {
return nil, nil
}
ak := strings.TrimSpace(input.AccessKey)
sk := strings.TrimSpace(input.SecretKey)
if ak == "" || sk == "" {
return nil, errors.New("移动云素材 AccessKey 和 SecretKey 必须同时填写")
}
return &model.ChannelAssetCredential{AccessKey: ak, SecretKey: sk, PoolID: strings.TrimSpace(input.PoolID)}, nil
}

func getVertexArrayKeys(keys string) ([]string, error) {
@@ -729,6 +785,15 @@ func AddChannel(c *gin.Context) {
})
return
}
credential, err := channelAssetCredentialFromInput(addChannelRequest.Channel.Type, addChannelRequest.AssetCredential)
if err != nil {
c.JSON(http.StatusOK, gin.H{"success": false, "message": err.Error()})
return
}
if credential != nil && addChannelRequest.Mode == "batch" {
c.JSON(http.StatusOK, gin.H{"success": false, "message": "移动云素材凭证不支持批量添加渠道"})
return
}

addChannelRequest.Channel.CreatedTime = common.GetTimestamp()
keys := make([]string, 0)
@@ -800,7 +865,15 @@ func AddChannel(c *gin.Context) {
}
channels = append(channels, *localChannel)
}
err = model.BatchInsertChannels(channels)
if credential != nil {
if len(channels) != 1 {
c.JSON(http.StatusOK, gin.H{"success": false, "message": "移动云素材凭证仅支持创建一个渠道"})
return
}
err = model.InsertChannelWithAssetCredential(&channels[0], credential)
} else {
err = model.BatchInsertChannels(channels)
}
if err != nil {
common.ApiError(c, err)
return
@@ -985,8 +1058,9 @@ func DeleteChannelBatch(c *gin.Context) {

type PatchChannel struct {
model.Channel
MultiKeyMode *string `json:"multi_key_mode"`
KeyMode *string `json:"key_mode"` // 多key模式下密钥覆盖或者追加
MultiKeyMode *string `json:"multi_key_mode"`
KeyMode *string `json:"key_mode"` // 多key模式下密钥覆盖或者追加
AssetCredential *ChannelAssetCredentialInput `json:"asset_credential"`
}

func UpdateChannel(c *gin.Context) {
@@ -1103,7 +1177,22 @@ func UpdateChannel(c *gin.Context) {
// 覆盖模式:直接使用新密钥(默认行为,不需要特殊处理)
}
}
err = channel.Update()
credential, err := channelAssetCredentialFromInput(channel.Type, channel.AssetCredential)
if err != nil {
c.JSON(http.StatusOK, gin.H{"success": false, "message": err.Error()})
return
}
if credential != nil {
err = model.DB.Transaction(func(tx *gorm.DB) error {
if err := channel.UpdateWithTx(tx); err != nil {
return err
}
credential.ChannelId = channel.Id
return model.UpsertChannelAssetCredentialWithTx(tx, credential)
})
} else {
err = channel.Update()
}
if err != nil {
common.ApiError(c, err)
return
@@ -1112,6 +1201,10 @@ func UpdateChannel(c *gin.Context) {
service.ResetProxyClientCache()
channel.Key = ""
clearChannelInfo(&channel.Channel)
if err := attachChannelAssetCredentialSummaries([]*model.Channel{&channel.Channel}); err != nil {
common.ApiError(c, err)
return
}
c.JSON(http.StatusOK, gin.H{
"success": true,
"message": "",
@@ -2105,4 +2198,3 @@ func OllamaVersion(c *gin.Context) {
},
})
}


+ 35
- 0
controller/channel_affinity_cache_test.go Просмотреть файл

@@ -0,0 +1,35 @@
package controller

import (
"net/http"
"net/http/httptest"
"testing"

"github.com/gin-gonic/gin"
"github.com/stretchr/testify/require"
)

func TestGetChannelAffinityUsageCacheStatsRequiresRuleAndFingerprint(t *testing.T) {
gin.SetMode(gin.TestMode)
router := gin.New()
router.GET("/stats", GetChannelAffinityUsageCacheStats)

for _, target := range []string{"/stats?key_fp=abc", "/stats?rule_name=rule"} {
response := httptest.NewRecorder()
router.ServeHTTP(response, httptest.NewRequest(http.MethodGet, target, nil))
require.Equal(t, http.StatusBadRequest, response.Code)
require.Contains(t, response.Body.String(), `"success":false`)
}
}

func TestClearChannelAffinityCacheRequiresSelector(t *testing.T) {
gin.SetMode(gin.TestMode)
router := gin.New()
router.DELETE("/cache", ClearChannelAffinityCache)

response := httptest.NewRecorder()
router.ServeHTTP(response, httptest.NewRequest(http.MethodDelete, "/cache", nil))

require.Equal(t, http.StatusBadRequest, response.Code)
require.Contains(t, response.Body.String(), `"success":false`)
}

+ 50
- 0
controller/channel_asset_credential_test.go Просмотреть файл

@@ -0,0 +1,50 @@
package controller

import (
"testing"

"github.com/QuantumNous/new-api/constant"
"github.com/QuantumNous/new-api/model"
"github.com/glebarez/sqlite"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"gorm.io/gorm"
"gorm.io/gorm/logger"
)

func setupChannelAssetCredentialControllerDB(t *testing.T) *gorm.DB {
t.Helper()
db, err := gorm.Open(sqlite.Open("file:controller_channel_asset_credentials?mode=memory&cache=shared"), &gorm.Config{
Logger: logger.Default.LogMode(logger.Silent),
})
require.NoError(t, err)
sqlDB, err := db.DB()
require.NoError(t, err)
sqlDB.SetMaxOpenConns(1)
originalDB := model.DB
model.DB = db
require.NoError(t, db.AutoMigrate(&model.Channel{}, &model.ChannelAssetCredential{}))
t.Cleanup(func() {
model.DB = originalDB
require.NoError(t, sqlDB.Close())
})
return db
}

func TestAttachChannelAssetCredentialSummariesDoesNotExposeSecrets(t *testing.T) {
db := setupChannelAssetCredentialControllerDB(t)
channel := &model.Channel{Id: 61, Type: constant.ChannelTypeChinaMobileSeedance, Key: "video-key", Name: "channel"}
require.NoError(t, db.Create(channel).Error)
require.NoError(t, model.UpsertChannelAssetCredential(&model.ChannelAssetCredential{
ChannelId: 61,
AccessKey: "ak-secret",
SecretKey: "sk-secret",
PoolID: "pool-61",
}))

require.NoError(t, attachChannelAssetCredentialSummaries([]*model.Channel{channel}))
assert.True(t, channel.AssetCredentialConfigured)
assert.Equal(t, "pool-61", channel.AssetCredentialPoolID)
assert.NotContains(t, channel.Key, "ak-secret")
assert.NotContains(t, channel.Key, "sk-secret")
}

+ 25
- 0
controller/channel_billing_helpers_test.go Просмотреть файл

@@ -0,0 +1,25 @@
package controller

import (
"net/http"
"testing"

"github.com/stretchr/testify/require"
)

func TestGetAuthHeadersUseProviderSpecificNames(t *testing.T) {
openAI := GetAuthHeader("openai-key")
claude := GetClaudeAuthHeader("claude-key")

require.Equal(t, "Bearer openai-key", openAI.Get("Authorization"))
require.Empty(t, openAI.Get("x-api-key"))
require.Equal(t, "claude-key", claude.Get("x-api-key"))
require.Equal(t, "2023-06-01", claude.Get("anthropic-version"))
require.Empty(t, claude.Get("Authorization"))
}

func TestGetResponseBodyRejectsInvalidMethodURLBeforeNetwork(t *testing.T) {
_, err := GetResponseBody("GET", "://invalid", nil, http.Header{})

require.Error(t, err)
}

+ 45
- 0
controller/deployment_test.go Просмотреть файл

@@ -0,0 +1,45 @@
package controller

import (
"net/http/httptest"
"testing"
"time"

"github.com/QuantumNous/new-api/pkg/ionet"
"github.com/gin-gonic/gin"
"github.com/stretchr/testify/require"
)

func TestMapIoNetDeploymentNormalizesStatusAndRemainingTime(t *testing.T) {
createdAt := time.Date(2026, 7, 25, 12, 0, 0, 0, time.UTC)
mapped := mapIoNetDeployment(ionet.Deployment{
ID: "deployment-1", Name: "Model Host", Status: "RUNNING", CreatedAt: createdAt,
BrandName: "NVIDIA", HardwareName: "H100", HardwareQuantity: 2, ComputeMinutesRemaining: 125,
})

require.Equal(t, "running", mapped["status"])
require.Equal(t, "2 hour 5 minutes", mapped["time_remaining"])
require.Equal(t, "NVIDIA H100 x2", mapped["hardware_info"])
require.EqualValues(t, createdAt.Unix(), mapped["created_at"])
}

func TestComputeStatusCountsIncludesKnownAndUnknownStatuses(t *testing.T) {
counts := computeStatusCounts(4, []ionet.Deployment{
{Status: "RUNNING"}, {Status: "running"}, {Status: "custom"},
})

require.EqualValues(t, 4, counts["all"])
require.EqualValues(t, 2, counts["running"])
require.EqualValues(t, 0, counts["failed"])
require.EqualValues(t, 1, counts["custom"])
}

func TestRequireDeploymentAndContainerIDsRejectBlankParameters(t *testing.T) {
context, _ := gin.CreateTestContext(httptest.NewRecorder())
context.Params = gin.Params{{Key: "id", Value: " "}, {Key: "container_id", Value: ""}}

_, ok := requireDeploymentID(context)
require.False(t, ok)
_, ok = requireContainerID(context)
require.False(t, ok)
}

+ 74
- 120
controller/doubao_asset.go Просмотреть файл

@@ -5,30 +5,17 @@ import (
"fmt"
"io"
"net/http"
"net/url"
"strings"
"time"

"github.com/QuantumNous/new-api/common"
"github.com/QuantumNous/new-api/constant"
"github.com/QuantumNous/new-api/logger"
"github.com/QuantumNous/new-api/model"
"github.com/QuantumNous/new-api/service"
"github.com/QuantumNous/new-api/setting/system_setting"

"github.com/gin-gonic/gin"
)

const defaultDoubaoAssetBaseURL = "https://ark.cn-beijing.volcengineapi.com"

var blockedAssetActions = map[string]struct{}{
"createassetgroup": {},
"getassetgroup": {},
"listassetgroups": {},
"updateassetgroup": {},
"deleteassetgroup": {},
}

func assetProxyError(c *gin.Context, status int, errType, message string) {
c.JSON(status, gin.H{
"error": gin.H{
@@ -38,36 +25,7 @@ func assetProxyError(c *gin.Context, status int, errType, message string) {
})
}

func buildDoubaoAssetURL(channel *model.Channel, action string, version string) (string, error) {
baseURL := defaultDoubaoAssetBaseURL
if channel != nil && channel.BaseURL != nil {
if configuredBaseURL := strings.TrimSpace(*channel.BaseURL); configuredBaseURL != "" {
baseURL = configuredBaseURL
}
}

u, err := url.Parse(baseURL)
if err != nil {
return "", err
}
if u.Scheme == "" || u.Host == "" {
return "", fmt.Errorf("invalid Doubao asset base URL: %s", baseURL)
}

u.Path = strings.TrimRight(u.Path, "/") + "/api/v1/multimodal/sd/assets"
u.RawQuery = ""
u.Fragment = ""

query := u.Query()
query.Set("Action", action)
if strings.TrimSpace(version) == "" {
version = "2024-01-01"
}
query.Set("Version", version)
u.RawQuery = query.Encode()

return u.String(), nil
}
var doubaoAssetAutoGroups = service.GetUserAutoGroup

func effectiveDoubaoAssetGroup(c *gin.Context) string {
if group := strings.TrimSpace(common.GetContextKeyString(c, constant.ContextKeyUsingGroup)); group != "" {
@@ -98,34 +56,15 @@ func concreteDoubaoAssetGroupsForRequest(c *gin.Context, autoGroups func(string)
return concreteGroups
}

func resolveDoubaoAssetChannelForRequest(c *gin.Context) (*model.Channel, string, error) {
userId := c.GetInt("id")
var lastErr error
for _, group := range concreteDoubaoAssetGroupsForRequest(c, service.GetUserAutoGroup) {
channel, err := service.ResolveDoubaoAssetChannel(userId, group)
if err != nil {
lastErr = err
logger.LogError(c.Request.Context(), fmt.Sprintf("Failed to resolve Doubao asset channel for group %s: %v", group, err))
continue
}
if channel != nil {
return channel, group, nil
}
}
if lastErr != nil {
return nil, "", lastErr
}
return nil, "", nil
}

func DoubaoAssetProxy(c *gin.Context) {
action := strings.TrimSpace(c.Query("Action"))
if action == "" {
assetProxyError(c, http.StatusBadRequest, "invalid_request_error", "Action query parameter is required")
actionRaw := strings.TrimSpace(c.Query("Action"))
if actionRaw == "" {
assetProxyError(c, http.StatusBadRequest, service.AssetErrorInvalidRequest, "Action query parameter is required")
return
}
if _, ok := blockedAssetActions[strings.ToLower(action)]; ok {
assetProxyError(c, http.StatusBadRequest, "invalid_request_error", fmt.Sprintf("Asset group API (%s) is not supported", action))
action, ok := service.ParseAssetAction(actionRaw)
if !ok {
assetProxyError(c, http.StatusBadRequest, service.AssetErrorInvalidRequest, fmt.Sprintf("unsupported asset Action: %s", actionRaw))
return
}

@@ -134,70 +73,85 @@ func DoubaoAssetProxy(c *gin.Context) {
version = "2024-01-01"
}

channel, _, err := resolveDoubaoAssetChannelForRequest(c)
rawBody, err := io.ReadAll(c.Request.Body)
if err != nil {
assetProxyError(c, http.StatusBadGateway, "server_error", fmt.Sprintf("Failed to resolve Doubao asset channel: %v", err))
return
}
if channel == nil {
assetProxyError(c, http.StatusBadGateway, "server_error", "Failed to resolve Doubao asset channel")
assetProxyError(c, http.StatusBadRequest, service.AssetErrorInvalidRequest, fmt.Sprintf("failed to read request body: %v", err))
return
}
if strings.TrimSpace(channel.Key) == "" {
assetProxyError(c, http.StatusBadGateway, "server_error", "Doubao asset channel key is missing")
return
}

upstreamURL, err := buildDoubaoAssetURL(channel, action, version)
if err != nil {
assetProxyError(c, http.StatusBadGateway, "server_error", fmt.Sprintf("Failed to build upstream URL: %v", err))
return
}

fetchSetting := system_setting.GetFetchSetting()
if err := common.ValidateURLWithFetchSetting(upstreamURL, fetchSetting.EnableSSRFProtection, fetchSetting.AllowPrivateIp, fetchSetting.DomainFilterMode, fetchSetting.IpFilterMode, fetchSetting.DomainList, fetchSetting.IpList, fetchSetting.AllowedPorts, fetchSetting.ApplyIPFilterForDomain); err != nil {
logger.LogError(c.Request.Context(), fmt.Sprintf("Doubao asset URL blocked: %v", err))
assetProxyError(c, http.StatusForbidden, "server_error", fmt.Sprintf("request blocked: %v", err))
return
}

client, err := service.GetHttpClientWithProxy(channel.GetSetting().Proxy)
if err != nil {
assetProxyError(c, http.StatusBadGateway, "server_error", fmt.Sprintf("Failed to create proxy client: %v", err))
return
}
if client == nil {
assetProxyError(c, http.StatusBadGateway, "server_error", "Failed to create proxy client")
return
body := map[string]any{}
if len(strings.TrimSpace(string(rawBody))) > 0 {
if err := common.Unmarshal(rawBody, &body); err != nil {
assetProxyError(c, http.StatusBadRequest, service.AssetErrorInvalidRequest, fmt.Sprintf("invalid JSON body: %v", err))
return
}
}

ctx, cancel := context.WithTimeout(c.Request.Context(), 60*time.Second)
defer cancel()

req, err := http.NewRequestWithContext(ctx, http.MethodPost, upstreamURL, c.Request.Body)
if err != nil {
assetProxyError(c, http.StatusBadGateway, "server_error", fmt.Sprintf("Failed to create upstream request: %v", err))
userID := c.GetInt("id")
var lastAssetErr *service.AssetError
for _, group := range concreteDoubaoAssetGroupsForRequest(c, doubaoAssetAutoGroups) {
channel, adapter, assetErr := service.ResolveAssetChannelForOperation(userID, group, action.Operation)
if assetErr != nil {
lastAssetErr = assetErr
logger.LogError(c.Request.Context(), fmt.Sprintf("Failed to resolve asset channel for group %s action %s: %s", group, action.Action, assetErr.Message))
if assetErr.Type == service.AssetErrorChannelNotFound {
continue
}
break
}
if channel == nil || adapter == nil {
continue
}
assetRequest := service.AssetRequest{
Action: action,
Version: version,
Body: body,
RawBody: rawBody,
}
if service.IsChinaMobileAssetChannel(channel) {
if service.IsAssetGroupOperation(action.Operation) {
assetProxyError(c, http.StatusForbidden, service.AssetErrorOperationNotSupported, "China Mobile asset group APIs are managed by the platform")
return
}
groupID, assetErr := service.GetOrCreateChinaMobileUserAssetGroup(c.Request.Context(), userID, channel, func(ctx context.Context, userID int, channel *model.Channel) (string, *service.AssetError) {
return service.CreateChinaMobileUserAssetGroup(ctx, userID, adapter, channel)
})
if assetErr != nil {
assetProxyError(c, assetErr.HTTPStatus, assetErr.Type, assetErr.Message)
return
}
service.ScopeChinaMobileAssetRequest(&assetRequest, groupID)
if action.Operation == service.AssetOperationAssetGet || action.Operation == service.AssetOperationAssetUpdate || action.Operation == service.AssetOperationAssetDelete {
if assetErr := service.RequireChinaMobileAssetOwnership(c.Request.Context(), adapter, channel, assetRequest, groupID); assetErr != nil {
assetProxyError(c, assetErr.HTTPStatus, assetErr.Type, assetErr.Message)
return
}
}
}
resp, assetErr := adapter.DoAssetRequest(c.Request.Context(), channel, assetRequest)
if assetErr != nil {
logger.LogError(c.Request.Context(), fmt.Sprintf("Failed to proxy asset action %s via channel %d type %d: %s", action.Action, channel.Id, channel.Type, assetErr.Message))
assetProxyError(c, assetErr.HTTPStatus, assetErr.Type, assetErr.Message)
return
}
copyDoubaoAssetResponseHeaders(c, resp.Header)
c.Data(resp.StatusCode, "application/json", resp.Body)
return
}
req.Header.Set("Content-Type", "application/json")
req.Header.Set("Accept", "application/json")
req.Header.Set("Authorization", "Bearer "+strings.TrimSpace(channel.Key))

resp, err := client.Do(req)
if err != nil {
logger.LogError(c.Request.Context(), fmt.Sprintf("Failed to proxy Doubao asset request to %s: %s", upstreamURL, err.Error()))
assetProxyError(c, http.StatusBadGateway, "server_error", fmt.Sprintf("Failed to proxy Doubao asset request: %v", err))
if lastAssetErr != nil {
assetProxyError(c, lastAssetErr.HTTPStatus, lastAssetErr.Type, lastAssetErr.Message)
return
}
defer resp.Body.Close()
assetProxyError(c, http.StatusBadGateway, service.AssetErrorChannelNotFound, "no available asset channel supports requested operation")
}

for key, values := range resp.Header {
func copyDoubaoAssetResponseHeaders(c *gin.Context, headers http.Header) {
for key, values := range headers {
switch http.CanonicalHeaderKey(key) {
case "Connection", "Content-Length", "Keep-Alive", "Proxy-Authenticate", "Proxy-Authorization", "Te", "Trailer", "Transfer-Encoding", "Upgrade":
continue
}
for _, value := range values {
c.Writer.Header().Add(key, value)
}
}
c.Writer.WriteHeader(resp.StatusCode)
if _, err = io.Copy(c.Writer, resp.Body); err != nil {
logger.LogError(c.Request.Context(), fmt.Sprintf("Failed to copy Doubao asset upstream response: %s", err.Error()))
}
}

+ 189
- 0
controller/doubao_asset_real_e2e_test.go Просмотреть файл

@@ -0,0 +1,189 @@
package controller

import (
"bytes"
"net/http"
"net/http/httptest"
"os"
"strings"
"testing"
"time"

"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/require"
"gorm.io/gorm"
"gorm.io/gorm/logger"
)

const chinaMobileAssetE2EImageURL = "https://bkimg.cdn.bcebos.com/pic/caef76094b36acaf2edd2133a78e9a1001e9380136fe?x-bce-process=image/format,f_auto/watermark,image_d2F0ZXIvYmFpa2UyNzI,g_7,xp_5,yp_5,P_20/resize,m_lfit,limit_1,h_1080"

func TestChinaMobileAssetControllerWithRealUpstream(t *testing.T) {
ak := strings.TrimSpace(os.Getenv("CHINAMOBILE_ASSET_AK"))
sk := strings.TrimSpace(os.Getenv("CHINAMOBILE_ASSET_SK"))
poolID := strings.TrimSpace(os.Getenv("CHINAMOBILE_ASSET_POOL_ID"))
if ak == "" || sk == "" {
t.Skip("CHINAMOBILE_ASSET_AK and CHINAMOBILE_ASSET_SK are required")
}
if poolID == "" {
poolID = "CIDC-CORE-00"
}
db := setupRealAssetE2EDB(t)
createDoubaoAssetProxyChannel(t, db, 7101, constant.ChannelTypeChinaMobileSeedance, "default", "video-generation-key", common.ChannelStatusEnabled)
require.NoError(t, model.UpsertChannelAssetCredential(&model.ChannelAssetCredential{
ChannelId: 7101,
AccessKey: ak,
SecretKey: sk,
PoolID: poolID,
}))
router := setupRealAssetE2ERouter()
suffix := time.Now().Format("20060102150405")
groupName := "new-api-e2e-" + suffix
assetName := "new-api-img-" + suffix

var groupId string
var assetId string
defer func() {
if assetId != "" {
code, _ := performRealAssetE2EAction(t, router, "DeleteAsset", map[string]any{"Id": assetId})
require.Equal(t, http.StatusOK, code, "cleanup DeleteAsset failed")
}
if groupId != "" {
code, _ := performRealAssetE2EAction(t, router, "DeleteAssetGroup", map[string]any{"Id": groupId})
require.Equal(t, http.StatusOK, code, "cleanup DeleteAssetGroup failed")
}
}()

code, body := performRealAssetE2EAction(t, router, "CreateAssetGroup", map[string]any{
"Name": groupName,
"GroupType": "AIGC",
"Description": "new-api real e2e test",
})
require.Equal(t, http.StatusOK, code)
groupId = realAssetE2EString(t, body, "Result.GroupId")
require.NotEmpty(t, groupId)

code, body = performRealAssetE2EAction(t, router, "CreateAsset", map[string]any{
"GroupId": groupId,
"Name": assetName,
"URL": chinaMobileAssetE2EImageURL,
"AssetType": "Image",
})
require.Equal(t, http.StatusOK, code)
assetId = realAssetE2EString(t, body, "Result")
require.NotEmpty(t, assetId)

var status string
for i := 0; i < 6; i++ {
code, body = performRealAssetE2EAction(t, router, "GetAsset", map[string]any{"Id": assetId})
require.Equal(t, http.StatusOK, code)
require.Equal(t, assetId, realAssetE2EString(t, body, "Result.Id"))
status = realAssetE2EString(t, body, "Result.Status")
if status != "Processing" {
break
}
time.Sleep(3 * time.Second)
}
require.Equal(t, "Active", status)

code, body = performRealAssetE2EAction(t, router, "ListAssets", map[string]any{
"PageNumber": 1,
"PageSize": 10,
"Filter": map[string]any{
"GroupType": "AIGC",
"GroupIds": []string{groupId},
},
})
require.Equal(t, http.StatusOK, code)
require.Equal(t, float64(1), realAssetE2EValue(t, body, "Result.TotalCount"))
}

func setupRealAssetE2EDB(t *testing.T) *gorm.DB {
t.Helper()
oldDB := model.DB
oldSQLitePath := common.SQLitePath
oldMemoryCacheEnabled := common.MemoryCacheEnabled
oldIsMasterNode := common.IsMasterNode
oldUsingSQLite := common.UsingSQLite
oldUsingMySQL := common.UsingMySQL
oldUsingPostgreSQL := common.UsingPostgreSQL
oldSQLDSN, hadSQLDSN := os.LookupEnv("SQL_DSN")

common.SQLitePath = "file:" + strings.NewReplacer("/", "_", " ", "_").Replace(t.Name()) + "?mode=memory&cache=shared"
common.MemoryCacheEnabled = 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())

model.DB = model.DB.Session(&gorm.Session{Logger: logger.Default.LogMode(logger.Silent)})
db := model.DB
sqlDB, err := db.DB()
require.NoError(t, err)
sqlDB.SetMaxOpenConns(1)
require.NoError(t, db.AutoMigrate(&model.Channel{}, &model.Ability{}, &model.UserAssetChannel{}, &model.ChannelAssetCredential{}))

t.Cleanup(func() {
_ = sqlDB.Close()
model.DB = oldDB
common.SQLitePath = oldSQLitePath
common.MemoryCacheEnabled = oldMemoryCacheEnabled
common.IsMasterNode = oldIsMasterNode
common.UsingSQLite = oldUsingSQLite
common.UsingMySQL = oldUsingMySQL
common.UsingPostgreSQL = oldUsingPostgreSQL
if hadSQLDSN {
_ = os.Setenv("SQL_DSN", oldSQLDSN)
} else {
_ = os.Unsetenv("SQL_DSN")
}
})
return db
}

func setupRealAssetE2ERouter() *gin.Engine {
router := gin.New()
router.Use(func(c *gin.Context) {
c.Set("id", 10)
common.SetContextKey(c, constant.ContextKeyUsingGroup, "default")
c.Next()
})
router.POST("/api/v1/volcengine/asset", DoubaoAssetProxy)
return router
}

func performRealAssetE2EAction(t *testing.T, router *gin.Engine, action string, payload map[string]any) (int, string) {
t.Helper()
data, err := common.Marshal(payload)
require.NoError(t, err)
req := httptest.NewRequest(http.MethodPost, "/api/v1/volcengine/asset?Action="+action, bytes.NewReader(data))
req.Header.Set("Content-Type", "application/json")
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
return w.Code, w.Body.String()
}

func realAssetE2EString(t *testing.T, data string, path string) string {
t.Helper()
value := realAssetE2EValue(t, data, path)
s, ok := value.(string)
require.Truef(t, ok, "%s is not a string: %#v", path, value)
return s
}

func realAssetE2EValue(t *testing.T, data string, path string) any {
t.Helper()
var payload map[string]any
require.NoError(t, common.Unmarshal([]byte(data), &payload))
current := any(payload)
for _, part := range strings.Split(path, ".") {
m, ok := current.(map[string]any)
require.Truef(t, ok, "%s is not an object at %s", path, part)
current = m[part]
}
return current
}

+ 330
- 14
controller/doubao_asset_test.go Просмотреть файл

@@ -1,16 +1,23 @@
package controller

import (
"context"
"fmt"
"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/QuantumNous/new-api/service"
"github.com/gin-gonic/gin"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"gorm.io/gorm"
"gorm.io/gorm/logger"
)

func setupDoubaoAssetProxyRouter(t *testing.T) *gin.Engine {
@@ -40,6 +47,74 @@ func decodeDoubaoAssetErrorMessage(t *testing.T, body string) string {
return payload.Error.Message
}

func setupDoubaoAssetProxyDB(t *testing.T) *gorm.DB {
t.Helper()

oldDB := model.DB
oldSQLitePath := common.SQLitePath
oldMemoryCacheEnabled := common.MemoryCacheEnabled
oldIsMasterNode := common.IsMasterNode
oldUsingSQLite := common.UsingSQLite
oldUsingMySQL := common.UsingMySQL
oldUsingPostgreSQL := common.UsingPostgreSQL
oldSQLDSN, hadSQLDSN := os.LookupEnv("SQL_DSN")

common.SQLitePath = "file:" + strings.NewReplacer("/", "_", " ", "_").Replace(t.Name()) + "?mode=memory&cache=shared"
common.MemoryCacheEnabled = 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())

model.DB = model.DB.Session(&gorm.Session{Logger: logger.Default.LogMode(logger.Silent)})
db := model.DB
sqlDB, err := db.DB()
require.NoError(t, err)
sqlDB.SetMaxOpenConns(1)
require.NoError(t, db.AutoMigrate(&model.Channel{}, &model.Ability{}, &model.UserAssetChannel{}, &model.UserAssetGroup{}))

t.Cleanup(func() {
_ = sqlDB.Close()
model.DB = oldDB
common.SQLitePath = oldSQLitePath
common.MemoryCacheEnabled = oldMemoryCacheEnabled
common.IsMasterNode = oldIsMasterNode
common.UsingSQLite = oldUsingSQLite
common.UsingMySQL = oldUsingMySQL
common.UsingPostgreSQL = oldUsingPostgreSQL
if hadSQLDSN {
_ = os.Setenv("SQL_DSN", oldSQLDSN)
} else {
_ = os.Unsetenv("SQL_DSN")
}
})

return db
}

func createDoubaoAssetProxyChannel(t *testing.T, db *gorm.DB, id int, channelType int, group string, key string, status int) {
t.Helper()

priority := int64(id)
weight := uint(10)
autoBan := 1
require.NoError(t, db.Create(&model.Channel{
Id: id,
Type: channelType,
Key: key,
Status: status,
Name: fmt.Sprintf("channel-%d", id),
Group: group,
Models: "seedance-2",
Priority: &priority,
Weight: &weight,
AutoBan: &autoBan,
CreatedTime: int64(id),
}).Error)
}

func TestDoubaoAssetProxyMissingActionReturns400(t *testing.T) {
router := setupDoubaoAssetProxyRouter(t)

@@ -51,35 +126,230 @@ func TestDoubaoAssetProxyMissingActionReturns400(t *testing.T) {
assert.Equal(t, "Action query parameter is required", decodeDoubaoAssetErrorMessage(t, w.Body.String()))
}

func TestDoubaoAssetProxyBlocksAssetGroupActionsCaseInsensitively(t *testing.T) {
func TestDoubaoAssetProxyUnknownActionReturns400(t *testing.T) {
router := setupDoubaoAssetProxyRouter(t)

req := httptest.NewRequest(http.MethodPost, "/api/v1/volcengine/asset?Action=CreateRealPersonAuthSession", nil)
w := httptest.NewRecorder()
router.ServeHTTP(w, req)

assert.Equal(t, http.StatusBadRequest, w.Code)
assert.Contains(t, decodeDoubaoAssetErrorMessage(t, w.Body.String()), "unsupported asset Action")
}

func TestDoubaoAssetProxyInvalidJSONReturns400(t *testing.T) {
router := setupDoubaoAssetProxyRouter(t)

for _, action := range []string{"CreateAssetGroup", "getassetgroup", "LISTASSETGROUPS", "UpdateAssetGroup", "deleteassetgroup"} {
req := httptest.NewRequest(http.MethodPost, "/api/v1/volcengine/asset?Action=ListAssets", strings.NewReader(`{"PageNumber":`))
w := httptest.NewRecorder()
router.ServeHTTP(w, req)

assert.Equal(t, http.StatusBadRequest, w.Code)
assert.Contains(t, decodeDoubaoAssetErrorMessage(t, w.Body.String()), "invalid JSON body")
}

func TestDoubaoAssetProxyRejectsChinaMobileAssetGroupActions(t *testing.T) {
db := setupDoubaoAssetProxyDB(t)
fake := &fakeDoubaoAssetAdapter{}
t.Cleanup(service.OverrideAssetAdapterForTest(constant.ChannelTypeChinaMobileSeedance, fake))
createDoubaoAssetProxyChannel(t, db, 61, constant.ChannelTypeChinaMobileSeedance, "default", "video-generation-key", common.ChannelStatusEnabled)
router := gin.New()
router.Use(func(c *gin.Context) {
c.Set("id", 10)
common.SetContextKey(c, constant.ContextKeyUsingGroup, "default")
c.Next()
})
router.POST("/api/v1/volcengine/asset", DoubaoAssetProxy)

for _, action := range []string{"CreateAssetGroup", "ListAssetGroups", "GetAssetGroup", "UpdateAssetGroup", "DeleteAssetGroup"} {
t.Run(action, func(t *testing.T) {
req := httptest.NewRequest(http.MethodPost, "/api/v1/volcengine/asset?Action="+action, nil)
req := httptest.NewRequest(http.MethodPost, "/api/v1/volcengine/asset?Action="+action, strings.NewReader(`{"Id":"group-1","GroupType":"AIGC","Name":"g"}`))
w := httptest.NewRecorder()
req.Header.Set("Content-Type", "application/json")
router.ServeHTTP(w, req)

assert.Equal(t, http.StatusBadRequest, w.Code)
assert.Contains(t, decodeDoubaoAssetErrorMessage(t, w.Body.String()), "Asset group API ("+action+") is not supported")
assert.Equal(t, http.StatusForbidden, w.Code)
assert.Equal(t, 0, fake.calls)
})
}
}

func TestBuildDoubaoAssetURLDefaultBaseAndEscapedQuery(t *testing.T) {
got, err := buildDoubaoAssetURL(&model.Channel{}, "ApplyUploadInner&Space", "")
func TestDoubaoAssetProxyDoesNotBlockNonChinaMobileAssetGroupAction(t *testing.T) {
db := setupDoubaoAssetProxyDB(t)
fake := &fakeDoubaoAssetAdapter{}
t.Cleanup(service.OverrideAssetAdapterForTest(constant.ChannelTypeDoubaoVideoCompatibleAiping, fake))
createDoubaoAssetProxyChannel(t, db, 63, constant.ChannelTypeDoubaoVideoCompatibleAiping, "default", "video-generation-key", common.ChannelStatusEnabled)
router := newDoubaoAssetProxyTestRouter()

require.NoError(t, err)
assert.Equal(t, defaultDoubaoAssetBaseURL+"/api/v1/multimodal/sd/assets?Action=ApplyUploadInner%26Space&Version=2024-01-01", got)
req := httptest.NewRequest(http.MethodPost, "/api/v1/volcengine/asset?Action=ListAssetGroups", strings.NewReader(`{"PageNumber":1,"PageSize":10}`))
w := httptest.NewRecorder()
router.ServeHTTP(w, req)

assert.Equal(t, http.StatusOK, w.Code)
assert.Equal(t, service.AssetOperationAssetGroupList, fake.operation)
}

func TestBuildDoubaoAssetURLExplicitBaseURLOverridesDefault(t *testing.T) {
baseURL := "https://example.com/custom/"
func TestDoubaoAssetProxyAutoGroupContinuesAfterMissingGroup(t *testing.T) {
db := setupDoubaoAssetProxyDB(t)
fake := &fakeDoubaoAssetAdapter{}
t.Cleanup(service.OverrideAssetAdapterForTest(constant.ChannelTypeChinaMobileSeedance, fake))
createDoubaoAssetProxyChannel(t, db, 62, constant.ChannelTypeChinaMobileSeedance, "vip", "video-generation-key", common.ChannelStatusEnabled)
router := gin.New()
router.Use(func(c *gin.Context) {
c.Set("id", 10)
common.SetContextKey(c, constant.ContextKeyUsingGroup, "auto")
common.SetContextKey(c, constant.ContextKeyUserGroup, "default")
c.Next()
})
router.POST("/api/v1/volcengine/asset", DoubaoAssetProxy)
oldAutoGroups := doubaoAssetAutoGroups
doubaoAssetAutoGroups = func(string) []string {
return []string{"default", "vip"}
}
t.Cleanup(func() {
doubaoAssetAutoGroups = oldAutoGroups
})

got, err := buildDoubaoAssetURL(&model.Channel{BaseURL: &baseURL}, "CommitUploadInner", "2025-02-03")
req := httptest.NewRequest(http.MethodPost, "/api/v1/volcengine/asset?Action=ListAssets", strings.NewReader(`{"Filter":{"GroupType":"AIGC"},"PageNumber":1,"PageSize":10}`))
w := httptest.NewRecorder()
req.Header.Set("Content-Type", "application/json")
router.ServeHTTP(w, req)

require.NoError(t, err)
assert.Equal(t, "https://example.com/custom/api/v1/multimodal/sd/assets?Action=CommitUploadInner&Version=2025-02-03", got)
assert.Equal(t, http.StatusOK, w.Code)
assert.Equal(t, service.AssetOperationAssetList, fake.operation)
}

func TestDoubaoAssetProxyForwardsAdapterResponseHeaders(t *testing.T) {
db := setupDoubaoAssetProxyDB(t)
fake := &fakeDoubaoAssetAdapter{responseHeader: http.Header{"X-Upstream-Request-Id": []string{"upstream-1"}, "Content-Length": []string{"999"}}}
t.Cleanup(service.OverrideAssetAdapterForTest(constant.ChannelTypeChinaMobileSeedance, fake))
createDoubaoAssetProxyChannel(t, db, 62, constant.ChannelTypeChinaMobileSeedance, "default", "video-generation-key", common.ChannelStatusEnabled)
router := gin.New()
router.Use(func(c *gin.Context) {
c.Set("id", 10)
common.SetContextKey(c, constant.ContextKeyUsingGroup, "default")
c.Next()
})
router.POST("/api/v1/volcengine/asset", DoubaoAssetProxy)

req := httptest.NewRequest(http.MethodPost, "/api/v1/volcengine/asset?Action=ListAssets", strings.NewReader(`{"Filter":{"GroupType":"AIGC"},"PageNumber":1,"PageSize":10}`))
w := httptest.NewRecorder()
req.Header.Set("Content-Type", "application/json")
router.ServeHTTP(w, req)

assert.Equal(t, http.StatusOK, w.Code)
assert.Equal(t, "upstream-1", w.Header().Get("X-Upstream-Request-Id"))
assert.NotEqual(t, "999", w.Header().Get("Content-Length"))
}

func TestDoubaoAssetProxyScopesChinaMobileCreateAndListRequests(t *testing.T) {
db := setupDoubaoAssetProxyDB(t)
fake := &fakeDoubaoAssetAdapter{}
t.Cleanup(service.OverrideAssetAdapterForTest(constant.ChannelTypeChinaMobileSeedance, fake))
createDoubaoAssetProxyChannel(t, db, 71, constant.ChannelTypeChinaMobileSeedance, "default", "video-generation-key", common.ChannelStatusEnabled)
require.NoError(t, model.CreateUserAssetGroup(10, 71, "owned"))
router := newDoubaoAssetProxyTestRouter()

for _, tc := range []struct {
action string
body string
}{
{"CreateAsset", `{"GroupId":"forged","Name":"n","URL":"https://example.com/a.png","AssetType":"Image"}`},
{"ListAssets", `{"Filter":{"GroupType":"AIGC","GroupIds":["forged"]},"PageNumber":1,"PageSize":10}`},
} {
t.Run(tc.action, func(t *testing.T) {
fake.requests = nil
req := httptest.NewRequest(http.MethodPost, "/api/v1/volcengine/asset?Action="+tc.action, strings.NewReader(tc.body))
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
require.Equal(t, http.StatusOK, w.Code)
require.Len(t, fake.requests, 1)
if tc.action == "CreateAsset" {
assert.Equal(t, "owned", fake.requests[0].Body["GroupId"])
} else {
assert.Equal(t, []string{"owned"}, fake.requests[0].Body["Filter"].(map[string]any)["GroupIds"])
}
})
}
}

func TestDoubaoAssetProxyRejectsChinaMobileListAssetsWithoutFilter(t *testing.T) {
db := setupDoubaoAssetProxyDB(t)
createDoubaoAssetProxyChannel(t, db, 73, constant.ChannelTypeChinaMobileSeedance, "default", "video-generation-key", common.ChannelStatusEnabled)
require.NoError(t, model.CreateUserAssetGroup(10, 73, "owned"))
router := newDoubaoAssetProxyTestRouter()

for _, body := range []string{`{"PageNumber":1,"PageSize":10}`, `{"Filter":null,"PageNumber":1,"PageSize":10}`} {
w := httptest.NewRecorder()
req := httptest.NewRequest(http.MethodPost, "/api/v1/volcengine/asset?Action=ListAssets", strings.NewReader(body))
router.ServeHTTP(w, req)

assert.Equal(t, http.StatusBadRequest, w.Code)
assert.Contains(t, decodeDoubaoAssetErrorMessage(t, w.Body.String()), "Filter.GroupType is required")
}
}

func TestDoubaoAssetProxyHidesChinaMobileAssetsOwnedByAnotherUser(t *testing.T) {
db := setupDoubaoAssetProxyDB(t)
fake := &fakeDoubaoAssetAdapter{getAssetGroupID: "other"}
t.Cleanup(service.OverrideAssetAdapterForTest(constant.ChannelTypeChinaMobileSeedance, fake))
createDoubaoAssetProxyChannel(t, db, 72, constant.ChannelTypeChinaMobileSeedance, "default", "video-generation-key", common.ChannelStatusEnabled)
require.NoError(t, model.CreateUserAssetGroup(10, 72, "owned"))
router := newDoubaoAssetProxyTestRouter()

for _, action := range []string{"GetAsset", "UpdateAsset", "DeleteAsset"} {
t.Run(action, func(t *testing.T) {
fake.requests = nil
req := httptest.NewRequest(http.MethodPost, "/api/v1/volcengine/asset?Action="+action, strings.NewReader(`{"Id":"asset-other","Name":"n"}`))
w := httptest.NewRecorder()
router.ServeHTTP(w, req)

assert.Equal(t, http.StatusNotFound, w.Code)
require.Len(t, fake.requests, 1)
assert.Equal(t, service.AssetOperationAssetGet, fake.requests[0].Action.Operation)
})
}
}

func newDoubaoAssetProxyTestRouter() *gin.Engine {
router := gin.New()
router.Use(func(c *gin.Context) {
c.Set("id", 10)
common.SetContextKey(c, constant.ContextKeyUsingGroup, "default")
c.Next()
})
router.POST("/api/v1/volcengine/asset", DoubaoAssetProxy)
return router
}

func TestDoubaoAssetProxyAutoGroupDoesNotBypassInvalidBinding(t *testing.T) {
db := setupDoubaoAssetProxyDB(t)
createDoubaoAssetProxyChannel(t, db, 1, constant.ChannelTypeChinaMobileSeedance, "other", "bad-key", common.ChannelStatusEnabled)
createDoubaoAssetProxyChannel(t, db, 2, constant.ChannelTypeChinaMobileSeedance, "vip", "video-generation-key", common.ChannelStatusEnabled)
require.NoError(t, model.BindUserAssetChannel(10, constant.ChannelTypeChinaMobileSeedance, "default", 1))
router := gin.New()
router.Use(func(c *gin.Context) {
c.Set("id", 10)
common.SetContextKey(c, constant.ContextKeyUsingGroup, "auto")
common.SetContextKey(c, constant.ContextKeyUserGroup, "default")
c.Next()
})
router.POST("/api/v1/volcengine/asset", DoubaoAssetProxy)
oldAutoGroups := doubaoAssetAutoGroups
doubaoAssetAutoGroups = func(string) []string {
return []string{"default", "vip"}
}
t.Cleanup(func() {
doubaoAssetAutoGroups = oldAutoGroups
})

req := httptest.NewRequest(http.MethodPost, "/api/v1/volcengine/asset?Action=ListAssets", strings.NewReader(`{"Filter":{"GroupType":"AIGC"},"PageNumber":1,"PageSize":10}`))
w := httptest.NewRecorder()
req.Header.Set("Content-Type", "application/json")
router.ServeHTTP(w, req)

assert.Equal(t, http.StatusBadGateway, w.Code)
assert.Contains(t, w.Body.String(), service.AssetErrorBindingInvalid)
}

func TestEffectiveDoubaoAssetGroupUsesUsingGroupBeforeBlankTokenGroup(t *testing.T) {
@@ -101,3 +371,49 @@ func TestConcreteDoubaoAssetGroupsForAutoUsesUserAutoGroups(t *testing.T) {

assert.Equal(t, []string{"default", "vip"}, groups)
}

type fakeDoubaoAssetAdapter struct {
operation service.AssetOperation
responseHeader http.Header
requests []service.AssetRequest
getAssetGroupID string
calls int
}

func (a *fakeDoubaoAssetAdapter) Name() string {
return "fake_asset"
}

func (a *fakeDoubaoAssetAdapter) Supports(operation service.AssetOperation) bool {
return true
}

func (a *fakeDoubaoAssetAdapter) DoAssetRequest(_ context.Context, _ *model.Channel, req service.AssetRequest) (*service.AssetUpstreamResponse, *service.AssetError) {
a.operation = req.Action.Operation
a.requests = append(a.requests, req)
a.calls++
if req.Action.Operation == service.AssetOperationAssetGroupCreate {
body, err := service.BuildAssetSuccessResponse(req.Action.Action, req.Version, map[string]any{"GroupId": "generated-group"})
if err != nil {
return nil, &service.AssetError{Type: service.AssetErrorServer, Message: err.Error(), HTTPStatus: http.StatusInternalServerError}
}
return &service.AssetUpstreamResponse{StatusCode: http.StatusOK, Header: a.responseHeader, Body: body}, nil
}
if req.Action.Operation == service.AssetOperationAssetGet {
body, err := service.BuildAssetSuccessResponse(req.Action.Action, req.Version, map[string]any{"Id": "asset-1", "GroupId": a.getAssetGroupID})
if err != nil {
return nil, &service.AssetError{Type: service.AssetErrorServer, Message: err.Error(), HTTPStatus: http.StatusInternalServerError}
}
return &service.AssetUpstreamResponse{StatusCode: http.StatusOK, Header: a.responseHeader, Body: body}, nil
}
body, err := service.BuildAssetSuccessResponse(req.Action.Action, req.Version, map[string]any{
"Items": []map[string]any{
{"Id": "group-1", "Name": "g", "GroupType": "AIGC"},
},
"TotalCount": 1,
})
if err != nil {
return nil, &service.AssetError{Type: service.AssetErrorServer, Message: err.Error(), HTTPStatus: http.StatusInternalServerError}
}
return &service.AssetUpstreamResponse{StatusCode: http.StatusOK, Header: a.responseHeader, Body: body}, nil
}

+ 18
- 0
controller/helpers_test.go Просмотреть файл

@@ -0,0 +1,18 @@
package controller

import (
"testing"

"github.com/stretchr/testify/require"
)

func TestBoolToString(t *testing.T) {
require.Equal(t, "true", boolToString(true))
require.Equal(t, "false", boolToString(false))
}

func TestGetLegalContentUsesEnglishOnlyForExactEnglishLanguage(t *testing.T) {
require.Equal(t, "English", getLegalContent("Chinese", "English", "en"))
require.Equal(t, "Chinese", getLegalContent("Chinese", "English", "zh"))
require.Equal(t, "Chinese", getLegalContent("Chinese", "English", "en-US"))
}

+ 16
- 0
controller/kling_aiping_native.go Просмотреть файл

@@ -173,6 +173,10 @@ func KlingAipingNativeProxy(c *gin.Context) {
c.JSON(http.StatusServiceUnavailable, gin.H{"code": 503, "message": err.Error()})
return
}
if err := persistKlingAipingProxyBindingIfNeeded(c.GetInt("id"), group, route.BillingModel, channel); err != nil {
c.JSON(http.StatusServiceUnavailable, gin.H{"code": 503, "message": err.Error()})
return
}
}
if setupErr := middleware.SetupContextForSelectedChannel(c, channel, route.BillingModel); setupErr != nil {
c.JSON(setupErr.StatusCode, gin.H{"code": setupErr.GetErrorCode(), "message": setupErr.Error()})
@@ -207,6 +211,18 @@ func resolveKlingAipingBoundChannelForProxy(c *gin.Context, group, billingModel
return channel
}

func persistKlingAipingProxyBindingIfNeeded(userId int, group, billingModel string, channel *model.Channel) error {
group = strings.TrimSpace(group)
billingModel = strings.TrimSpace(billingModel)
if group == "" || group == "auto" || billingModel == "" || channel == nil {
return nil
}
if !service.IsUsableVideoAssetChannelForFamily(channel, group, billingModel, service.VideoAssetFamilyKling) {
return nil
}
return service.BindVideoAssetChannel(userId, group, channel, service.VideoAssetFamilyKling)
}

func readJSONPayload(c *gin.Context) (map[string]any, error) {
body, err := common.GetBodyStorage(c)
if err != nil {


+ 113
- 0
controller/kling_aiping_native_test.go Просмотреть файл

@@ -4,6 +4,7 @@ import (
"io"
"net/http"
"net/http/httptest"
"os"
"strings"
"testing"

@@ -16,6 +17,8 @@ import (
"github.com/QuantumNous/new-api/service"
"github.com/gin-gonic/gin"
"github.com/stretchr/testify/require"
"gorm.io/gorm"
"gorm.io/gorm/logger"
)

func TestKlingAipingSubmitPreparationUsesModelNameBeforeModel(t *testing.T) {
@@ -172,6 +175,34 @@ func TestDoKlingAipingProxyRequestUsesSelectedContextKey(t *testing.T) {
require.Equal(t, "Bearer selected-key", gotAuth)
}

func TestKlingAipingNativeProxyPersistsFallbackBinding(t *testing.T) {
service.InitHttpClient()
db := setupKlingAipingNativeProxyDB(t)
upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
w.Header().Set("Content-Type", "application/json")
_, _ = w.Write([]byte(`{"code":0,"message":"success","data":[]}`))
}))
defer upstream.Close()
createKlingAipingProxyChannelForTest(t, db, 59, "default", "proxy-key", upstream.URL)
createKlingAipingProxyAbilityForTest(t, db, "default", klingaiping.ModelKlingAdvancedElements, 59, true)

w := httptest.NewRecorder()
_, engine := gin.CreateTestContext(w)
engine.GET("/v1/general/advanced-presets-elements", func(c *gin.Context) {
c.Set("id", 10)
common.SetContextKey(c, constant.ContextKeyUsingGroup, "default")
KlingAipingNativeProxy(c)
})

engine.ServeHTTP(w, httptest.NewRequest(http.MethodGet, "/v1/general/advanced-presets-elements", nil))

require.Equal(t, http.StatusOK, w.Code)
binding, err := model.GetUserAssetChannel(10, constant.ChannelTypeKlingAiping, "default")
require.NoError(t, err)
require.NotNil(t, binding)
require.Equal(t, 59, binding.ChannelId)
}

func TestCopyProxyResponseNormalizesMsgError(t *testing.T) {
w := httptest.NewRecorder()
c, _ := gin.CreateTestContext(w)
@@ -248,3 +279,85 @@ func TestParseKlingAipingPageBoundaries(t *testing.T) {
_, _, err = parseKlingAipingPage(c)
require.ErrorContains(t, err, "pageSize")
}

func setupKlingAipingNativeProxyDB(t *testing.T) *gorm.DB {
t.Helper()

oldDB := model.DB
oldSQLitePath := common.SQLitePath
oldMemoryCacheEnabled := common.MemoryCacheEnabled
oldIsMasterNode := common.IsMasterNode
oldUsingSQLite := common.UsingSQLite
oldUsingMySQL := common.UsingMySQL
oldUsingPostgreSQL := common.UsingPostgreSQL
oldSQLDSN, hadSQLDSN := os.LookupEnv("SQL_DSN")

common.SQLitePath = "file:" + strings.NewReplacer("/", "_", " ", "_").Replace(t.Name()) + "?mode=memory&cache=shared"
common.MemoryCacheEnabled = 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())

model.DB = model.DB.Session(&gorm.Session{Logger: logger.Default.LogMode(logger.Silent)})
db := model.DB
sqlDB, err := db.DB()
require.NoError(t, err)
sqlDB.SetMaxOpenConns(1)
require.NoError(t, db.AutoMigrate(&model.Channel{}, &model.Ability{}, &model.UserAssetChannel{}))

t.Cleanup(func() {
_ = sqlDB.Close()
model.DB = oldDB
common.SQLitePath = oldSQLitePath
common.MemoryCacheEnabled = oldMemoryCacheEnabled
common.IsMasterNode = oldIsMasterNode
common.UsingSQLite = oldUsingSQLite
common.UsingMySQL = oldUsingMySQL
common.UsingPostgreSQL = oldUsingPostgreSQL
if hadSQLDSN {
_ = os.Setenv("SQL_DSN", oldSQLDSN)
} else {
_ = os.Unsetenv("SQL_DSN")
}
})

return db
}

func createKlingAipingProxyChannelForTest(t *testing.T, db *gorm.DB, id int, group string, key string, baseURL string) {
t.Helper()

priority := int64(id)
weight := uint(10)
autoBan := 1
require.NoError(t, db.Create(&model.Channel{
Id: id,
Type: constant.ChannelTypeKlingAiping,
Key: key,
Status: common.ChannelStatusEnabled,
Name: "kling-aiping-proxy",
Group: group,
Models: klingaiping.ModelKlingAdvancedElements,
BaseURL: common.GetPointer(baseURL),
Priority: &priority,
Weight: &weight,
AutoBan: &autoBan,
}).Error)
}

func createKlingAipingProxyAbilityForTest(t *testing.T, db *gorm.DB, group string, modelName string, channelId int, enabled bool) {
t.Helper()

priority := int64(channelId)
require.NoError(t, db.Create(&model.Ability{
Group: group,
Model: modelName,
ChannelId: channelId,
Enabled: enabled,
Priority: &priority,
Weight: 10,
}).Error)
}

+ 40
- 0
controller/option_test.go Просмотреть файл

@@ -0,0 +1,40 @@
package controller

import (
"net/http"
"net/http/httptest"
"testing"

"github.com/QuantumNous/new-api/common"
"github.com/gin-gonic/gin"
"github.com/stretchr/testify/require"
)

func TestGetOptionsOmitsSensitiveSuffixes(t *testing.T) {
common.OptionMapRWMutex.Lock()
oldOptions := common.OptionMap
common.OptionMap = map[string]string{
"PublicOption": "visible", "AccessToken": "secret", "WebhookSecret": "secret",
"ProviderKey": "secret", "lowercase_secret": "secret", "service_api_key": "secret",
}
common.OptionMapRWMutex.Unlock()
t.Cleanup(func() {
common.OptionMapRWMutex.Lock()
common.OptionMap = oldOptions
common.OptionMapRWMutex.Unlock()
})

router := gin.New()
router.GET("/", GetOptions)
response := httptest.NewRecorder()
router.ServeHTTP(response, httptest.NewRequest(http.MethodGet, "/", nil))

require.Equal(t, http.StatusOK, response.Code)
body := response.Body.String()
require.Contains(t, body, "PublicOption")
require.NotContains(t, body, "AccessToken")
require.NotContains(t, body, "WebhookSecret")
require.NotContains(t, body, "ProviderKey")
require.NotContains(t, body, "lowercase_secret")
require.NotContains(t, body, "service_api_key")
}

+ 55
- 0
controller/payment_webhook_availability.go Просмотреть файл

@@ -0,0 +1,55 @@
package controller

import (
"strings"

"github.com/QuantumNous/new-api/setting"
"github.com/QuantumNous/new-api/setting/operation_setting"
)

// isStripeTopUpEnabled 检查 Stripe 支付是否已启用(三项配置缺一不可)
func isStripeTopUpEnabled() bool {
return strings.TrimSpace(setting.StripeApiSecret) != "" &&
strings.TrimSpace(setting.StripeWebhookSecret) != "" &&
strings.TrimSpace(setting.StripePriceId) != ""
}

// isStripeWebhookEnabled 检查 Stripe Webhook 是否可接收
func isStripeWebhookEnabled() bool {
return isStripeTopUpEnabled()
}

// isCreemTopUpEnabled 检查 Creem 支付是否已启用
func isCreemTopUpEnabled() bool {
products := strings.TrimSpace(setting.CreemProducts)
return strings.TrimSpace(setting.CreemApiKey) != "" &&
products != "" &&
products != "[]"
}

// isCreemWebhookEnabled 检查 Creem Webhook 是否可接收
func isCreemWebhookEnabled() bool {
return isCreemTopUpEnabled() && strings.TrimSpace(setting.CreemWebhookSecret) != ""
}

// isWechatPayWebhookEnabled 检查微信支付 Webhook 是否可接收
func isWechatPayWebhookEnabled() bool {
return setting.IsWechatPayConfigured()
}

// isAlipayWebhookEnabled 检查支付宝 Webhook 是否可接收
func isAlipayWebhookEnabled() bool {
return setting.IsAlipayConfigured()
}

// isEpayTopUpEnabled 检查易支付是否已启用
func isEpayTopUpEnabled() bool {
return strings.TrimSpace(operation_setting.PayAddress) != "" &&
strings.TrimSpace(operation_setting.EpayId) != "" &&
strings.TrimSpace(operation_setting.EpayKey) != ""
}

// isEpayWebhookEnabled 检查易支付 Webhook 是否可接收
func isEpayWebhookEnabled() bool {
return isEpayTopUpEnabled() && len(operation_setting.PayMethods) > 0
}

+ 57
- 0
controller/payment_webhook_availability_test.go Просмотреть файл

@@ -0,0 +1,57 @@
package controller

import (
"testing"

"github.com/QuantumNous/new-api/setting"
"github.com/QuantumNous/new-api/setting/operation_setting"
"github.com/stretchr/testify/require"
)

func TestStripeWebhookAvailabilityRequiresAllTopUpCredentials(t *testing.T) {
oldAPI, oldWebhook, oldPrice := setting.StripeApiSecret, setting.StripeWebhookSecret, setting.StripePriceId
t.Cleanup(func() {
setting.StripeApiSecret, setting.StripeWebhookSecret, setting.StripePriceId = oldAPI, oldWebhook, oldPrice
})

setting.StripeApiSecret, setting.StripeWebhookSecret, setting.StripePriceId = "api", "webhook", "price"
require.True(t, isStripeTopUpEnabled())
require.True(t, isStripeWebhookEnabled())

setting.StripeWebhookSecret = " "
require.False(t, isStripeTopUpEnabled())
require.False(t, isStripeWebhookEnabled())
}

func TestCreemWebhookAvailabilityRequiresProductsAndWebhookSecret(t *testing.T) {
oldAPI, oldProducts, oldWebhook := setting.CreemApiKey, setting.CreemProducts, setting.CreemWebhookSecret
t.Cleanup(func() {
setting.CreemApiKey, setting.CreemProducts, setting.CreemWebhookSecret = oldAPI, oldProducts, oldWebhook
})

setting.CreemApiKey, setting.CreemProducts, setting.CreemWebhookSecret = "api", "[]", "webhook"
require.False(t, isCreemTopUpEnabled())
require.False(t, isCreemWebhookEnabled())

setting.CreemProducts = `[{"id":"product"}]`
require.True(t, isCreemTopUpEnabled())
require.True(t, isCreemWebhookEnabled())

setting.CreemWebhookSecret = ""
require.False(t, isCreemWebhookEnabled())
}

func TestEpayWebhookAvailabilityRequiresPaymentMethod(t *testing.T) {
oldAddress, oldID, oldKey, oldMethods := operation_setting.PayAddress, operation_setting.EpayId, operation_setting.EpayKey, operation_setting.PayMethods
t.Cleanup(func() {
operation_setting.PayAddress, operation_setting.EpayId, operation_setting.EpayKey, operation_setting.PayMethods = oldAddress, oldID, oldKey, oldMethods
})

operation_setting.PayAddress, operation_setting.EpayId, operation_setting.EpayKey = "https://pay.example", "merchant", "secret"
operation_setting.PayMethods = nil
require.True(t, isEpayTopUpEnabled())
require.False(t, isEpayWebhookEnabled())

operation_setting.PayMethods = []map[string]string{{"name": "alipay"}}
require.True(t, isEpayWebhookEnabled())
}

+ 31
- 0
controller/ratio_config_test.go Просмотреть файл

@@ -0,0 +1,31 @@
package controller

import (
"net/http"
"net/http/httptest"
"testing"

"github.com/QuantumNous/new-api/setting/ratio_setting"
"github.com/gin-gonic/gin"
"github.com/stretchr/testify/require"
)

func TestGetRatioConfigRejectsRequestWhenExposureIsDisabled(t *testing.T) {
ratio_setting.SetExposeRatioEnabled(false)
t.Cleanup(func() { ratio_setting.SetExposeRatioEnabled(false) })
context, _ := gin.CreateTestContext(httptest.NewRecorder())

GetRatioConfig(context)

require.Equal(t, http.StatusForbidden, context.Writer.Status())
}

func TestGetRatioConfigReturnsDataWhenExposureIsEnabled(t *testing.T) {
ratio_setting.SetExposeRatioEnabled(true)
t.Cleanup(func() { ratio_setting.SetExposeRatioEnabled(false) })
context, _ := gin.CreateTestContext(httptest.NewRecorder())

GetRatioConfig(context)

require.Equal(t, http.StatusOK, context.Writer.Status())
}

+ 5
- 0
controller/subscription_payment_epay.go Просмотреть файл

@@ -112,6 +112,11 @@ func SubscriptionRequestEpay(c *gin.Context) {
}

func SubscriptionEpayNotify(c *gin.Context) {
if !isEpayWebhookEnabled() {
_, _ = c.Writer.Write([]byte("fail"))
return
}

var params map[string]string

if c.Request.Method == "POST" {


+ 8
- 1
controller/topup.go Просмотреть файл

@@ -258,7 +258,8 @@ func RequestEpay(c *gin.Context) {
Amount: amount,
Money: payMoney,
TradeNo: tradeNo,
PaymentMethod: req.PaymentMethod,
PaymentMethod: req.PaymentMethod,
PaymentProvider: model.PaymentProviderEpay,
CreateTime: time.Now().Unix(),
Status: "pending",
}
@@ -298,6 +299,12 @@ func UnlockOrder(tradeNo string) {
}

func EpayNotify(c *gin.Context) {
if !isEpayWebhookEnabled() {
log.Println("易支付 webhook 被拒绝: 易支付未配置或已禁用")
_, _ = c.Writer.Write([]byte("fail"))
return
}

var params map[string]string

if c.Request.Method == "POST" {


+ 8
- 1
controller/topup_alipay.go Просмотреть файл

@@ -190,7 +190,8 @@ func RequestAlipayPay(c *gin.Context) {
Amount: req.Amount,
Money: chargedMoney,
TradeNo: tradeNo,
PaymentMethod: PaymentMethodAlipay,
PaymentMethod: PaymentMethodAlipay,
PaymentProvider: model.PaymentProviderAlipay,
CreateTime: time.Now().Unix(),
Status: common.TopUpStatusPending,
}
@@ -239,6 +240,12 @@ func AlipayPayStatus(c *gin.Context) {

// AlipayPayWebhook 处理支付宝异步回调通知
func AlipayPayWebhook(c *gin.Context) {
if !isAlipayWebhookEnabled() {
log.Printf("支付宝 webhook 被拒绝: 支付宝未配置 (client_ip=%s)\n", c.ClientIP())
c.String(http.StatusForbidden, "fail")
return
}

notifyReq, err := alipay.ParseNotifyToBodyMap(c.Request)
if err != nil {
log.Printf("解析支付宝回调失败: %v", err)


+ 13
- 6
controller/topup_creem.go Просмотреть файл

@@ -108,12 +108,13 @@ func (*CreemAdaptor) RequestPay(c *gin.Context, req *CreemPayRequest) {

// 先创建订单记录,使用产品配置的金额和充值额度
topUp := &model.TopUp{
UserId: id,
Amount: selectedProduct.Quota, // 充值额度
Money: selectedProduct.Price, // 支付金额
TradeNo: referenceId,
CreateTime: time.Now().Unix(),
Status: common.TopUpStatusPending,
UserId: id,
Amount: selectedProduct.Quota, // 充值额度
Money: selectedProduct.Price, // 支付金额
TradeNo: referenceId,
PaymentProvider: model.PaymentProviderCreem,
CreateTime: time.Now().Unix(),
Status: common.TopUpStatusPending,
}
err = topUp.Insert()
if err != nil {
@@ -229,6 +230,12 @@ type CreemWebhookEvent struct {
}

func CreemWebhook(c *gin.Context) {
if !isCreemWebhookEnabled() {
log.Printf("Creem webhook 被拒绝: webhook 未配置或已禁用 (client_ip=%s)\n", c.ClientIP())
c.AbortWithStatus(http.StatusForbidden)
return
}

// 读取body内容用于打印,同时保留原始数据供后续使用
bodyBytes, err := io.ReadAll(c.Request.Body)
if err != nil {


+ 8
- 1
controller/topup_stripe.go Просмотреть файл

@@ -108,7 +108,8 @@ func (*StripeAdaptor) RequestPay(c *gin.Context, req *StripePayRequest) {
Amount: req.Amount,
Money: chargedMoney,
TradeNo: referenceId,
PaymentMethod: PaymentMethodStripe,
PaymentMethod: PaymentMethodStripe,
PaymentProvider: model.PaymentProviderStripe,
CreateTime: time.Now().Unix(),
Status: common.TopUpStatusPending,
}
@@ -146,6 +147,12 @@ func RequestStripePay(c *gin.Context) {
}

func StripeWebhook(c *gin.Context) {
if !isStripeWebhookEnabled() {
log.Printf("Stripe webhook 被拒绝: webhook 未配置或已禁用 (client_ip=%s)\n", c.ClientIP())
c.AbortWithStatus(http.StatusForbidden)
return
}

payload, err := io.ReadAll(c.Request.Body)
if err != nil {
log.Printf("解析Stripe Webhook参数失败: %v\n", err)


+ 8
- 1
controller/topup_wechat.go Просмотреть файл

@@ -246,7 +246,8 @@ func RequestWechatPay(c *gin.Context) {
Amount: req.Amount,
Money: chargedMoney,
TradeNo: tradeNo,
PaymentMethod: PaymentMethodWechatPay,
PaymentMethod: PaymentMethodWechatPay,
PaymentProvider: model.PaymentProviderWechat,
CreateTime: time.Now().Unix(),
Status: common.TopUpStatusPending,
}
@@ -296,6 +297,12 @@ func WechatPayStatus(c *gin.Context) {

// WechatPayWebhook 处理微信支付回调通知
func WechatPayWebhook(c *gin.Context) {
if !isWechatPayWebhookEnabled() {
log.Printf("微信支付 webhook 被拒绝: 微信支付未配置 (client_ip=%s)\n", c.ClientIP())
c.JSON(http.StatusForbidden, gin.H{"code": "FAIL", "message": "微信支付未配置"})
return
}

notifyReq, err := wechat.V3ParseNotify(c.Request)
if err != nil {
log.Printf("解析微信支付回调失败: %v", err)


+ 10
- 0
controller/user.go Просмотреть файл

@@ -837,6 +837,16 @@ func CreateUser(c *gin.Context) {
common.ApiErrorI18n(c, i18n.MsgUserCannotCreateHigherLevel)
return
}
exist, err := model.CheckUserExistOrDeleted(user.Username, user.Email)
if err != nil {
common.ApiErrorI18n(c, i18n.MsgDatabaseError)
common.SysLog(fmt.Sprintf("CheckUserExistOrDeleted error: %v", err))
return
}
if exist {
common.ApiErrorI18n(c, i18n.MsgUserExists)
return
}
// Even for admin users, we cannot fully trust them!
cleanUser := model.User{
Username: user.Username,


+ 209
- 0
controller/user_video_channel_binding.go Просмотреть файл

@@ -0,0 +1,209 @@
package controller

import (
"fmt"
"net/http"
"strconv"
"strings"

"github.com/QuantumNous/new-api/common"
"github.com/QuantumNous/new-api/model"
"github.com/QuantumNous/new-api/service"
"github.com/gin-gonic/gin"
"gorm.io/gorm"
)

type AdminVideoChannelBinding struct {
Group string `json:"group"`
Family string `json:"family"`
ChannelID int `json:"channel_id"`
}

type adminVideoChannelBindingRequest struct {
Bindings []AdminVideoChannelBinding `json:"bindings"`
}

type validatedAdminVideoChannelBinding struct {
AdminVideoChannelBinding
ChannelType int
}

type AdminVideoChannelCandidate struct {
ID int `json:"id"`
Name string `json:"name"`
Type int `json:"type"`
}

type adminVideoChannelBindingRow struct {
Group string `json:"group"`
ChannelID int `json:"channel_id"`
ChannelName string `json:"channel_name"`
ChannelType int `json:"channel_type"`
Candidates []AdminVideoChannelCandidate `json:"candidates"`
}

type adminVideoChannelBindingFamily struct {
Key string `json:"key"`
Name string `json:"name"`
Bindings []adminVideoChannelBindingRow `json:"bindings"`
}

func videoAssetFamilyFromString(value string) (service.VideoAssetFamily, bool) {
family := service.VideoAssetFamily(strings.TrimSpace(value))
for _, candidate := range service.VideoAssetFamilies() {
if family == candidate {
return family, true
}
}
return "", false
}

func GetUserVideoChannelBindings(c *gin.Context) {
userID, err := strconv.Atoi(c.Param("id"))
if err != nil || userID <= 0 {
common.ApiErrorMsg(c, "invalid user id")
return
}
if _, err = model.GetUserById(userID, false); err != nil {
common.ApiError(c, err)
return
}
groups, err := model.GetUserConcreteTokenGroups(userID)
if err != nil {
common.ApiError(c, err)
return
}
visibleGroupSet := make(map[string]struct{}, len(groups))
families := make([]adminVideoChannelBindingFamily, 0, len(service.VideoAssetFamilies()))
for _, family := range service.VideoAssetFamilies() {
rows := make([]adminVideoChannelBindingRow, 0, len(groups))
for _, group := range groups {
candidates, err := service.GetVideoAssetChannelCandidates(group, family)
if err != nil {
common.ApiError(c, err)
return
}
if len(candidates) == 0 {
continue
}
visibleGroupSet[group] = struct{}{}
row := adminVideoChannelBindingRow{Group: group, Candidates: make([]AdminVideoChannelCandidate, 0, len(candidates))}
for _, channel := range candidates {
row.Candidates = append(row.Candidates, AdminVideoChannelCandidate{ID: channel.Id, Name: channel.Name, Type: channel.Type})
}
bindings, err := model.GetUserAssetChannelsByTypes(userID, service.VideoAssetChannelTypesForFamily(family), group)
if err != nil {
common.ApiError(c, err)
return
}
for _, binding := range bindings {
for _, candidate := range candidates {
if candidate.Id == binding.ChannelId {
row.ChannelID = candidate.Id
row.ChannelName = candidate.Name
row.ChannelType = candidate.Type
break
}
}
if row.ChannelID != 0 {
break
}
}
rows = append(rows, row)
}
families = append(families, adminVideoChannelBindingFamily{Key: string(family), Name: strings.ToUpper(string(family[:1])) + string(family[1:]), Bindings: rows})
}
visibleGroups := make([]string, 0, len(visibleGroupSet))
for _, group := range groups {
if _, ok := visibleGroupSet[group]; ok {
visibleGroups = append(visibleGroups, group)
}
}
c.JSON(http.StatusOK, gin.H{"success": true, "message": "", "data": gin.H{"groups": visibleGroups, "families": families}})
}

func SetUserVideoChannelBindings(c *gin.Context) {
userID, err := strconv.Atoi(c.Param("id"))
if err != nil || userID <= 0 {
common.ApiErrorMsg(c, "invalid user id")
return
}
if _, err = model.GetUserById(userID, false); err != nil {
common.ApiError(c, err)
return
}
var request adminVideoChannelBindingRequest
if err := c.ShouldBindJSON(&request); err != nil {
common.ApiErrorMsg(c, err.Error())
return
}
groups, err := model.GetUserConcreteTokenGroups(userID)
if err != nil {
common.ApiError(c, err)
return
}
allowedGroups := make(map[string]struct{}, len(groups))
for _, group := range groups {
allowedGroups[group] = struct{}{}
}
requested := make(map[string]validatedAdminVideoChannelBinding, len(request.Bindings))
for _, binding := range request.Bindings {
binding.Group = strings.TrimSpace(binding.Group)
family, ok := videoAssetFamilyFromString(binding.Family)
if !ok || binding.Group == "" || binding.Group == "auto" || binding.ChannelID <= 0 {
common.ApiErrorMsg(c, "invalid video channel binding")
return
}
if _, ok := allowedGroups[binding.Group]; !ok {
common.ApiErrorMsg(c, "token group is not available for this user")
return
}
key := binding.Group + "\x00" + string(family)
if _, exists := requested[key]; exists {
common.ApiErrorMsg(c, "duplicate video channel binding")
return
}
candidates, err := service.GetVideoAssetChannelCandidates(binding.Group, family)
if err != nil {
common.ApiError(c, err)
return
}
channelType := 0
for _, candidate := range candidates {
if candidate.Id == binding.ChannelID {
channelType = candidate.Type
break
}
}
if channelType == 0 {
common.ApiErrorMsg(c, "channel is not available for this video family and token group")
return
}
binding.Family = string(family)
requested[key] = validatedAdminVideoChannelBinding{AdminVideoChannelBinding: binding, ChannelType: channelType}
}
if err := model.DB.Transaction(func(tx *gorm.DB) error {
for _, group := range groups {
for _, family := range service.VideoAssetFamilies() {
key := group + "\x00" + string(family)
binding, exists := requested[key]
channelTypes := service.VideoAssetChannelTypesForFamily(family)
if err := model.DeleteUserAssetChannelsByTypesWithTx(tx, userID, channelTypes, group); err != nil {
return err
}
if !exists {
continue
}
if err := model.BindUserAssetChannelWithTx(tx, userID, binding.ChannelType, group, binding.ChannelID); err != nil {
return err
}
}
}
return nil
}); err != nil {
common.ApiError(c, err)
return
}
model.RecordLog(userID, model.LogTypeManage, fmt.Sprintf("updated video channel bindings for user %d", userID))
c.JSON(http.StatusOK, gin.H{"success": true, "message": ""})
}

+ 243
- 0
controller/user_video_channel_binding_test.go Просмотреть файл

@@ -0,0 +1,243 @@
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)
}
}

+ 40
- 0
controller/video_proxy_gemini_test.go Просмотреть файл

@@ -0,0 +1,40 @@
package controller

import (
"testing"

"github.com/QuantumNous/new-api/common"
"github.com/QuantumNous/new-api/model"
"github.com/stretchr/testify/require"
)

func TestExtractGeminiVideoURLSupportsTopLevelAndNestedShapes(t *testing.T) {
require.Equal(t, "https://video.example/top", extractGeminiVideoURLFromMap(map[string]any{"uri": "https://video.example/top"}))

nested := map[string]any{
"response": map[string]any{
"generateVideoResponse": map[string]any{
"generatedSamples": []any{map[string]any{"video": map[string]any{"uri": "https://video.example/generated"}}},
},
},
}
require.Equal(t, "https://video.example/generated", extractGeminiVideoURLFromMap(nested))
require.Equal(t, "", extractGeminiVideoURLFromMap(map[string]any{"response": map[string]any{}}))
}

func TestExtractGeminiVideoURLFromTaskDataRejectsInvalidJSON(t *testing.T) {
task := &model.Task{Data: []byte("not-json")}
require.Empty(t, extractGeminiVideoURLFromTaskData(task))

payload, err := common.Marshal(map[string]any{"response": map[string]any{"video": "https://video.example/nested"}})
require.NoError(t, err)
task.Data = payload
require.Equal(t, "https://video.example/nested", extractGeminiVideoURLFromTaskData(task))
}

func TestEnsureAPIKeyAppendsOnlyWhenMissing(t *testing.T) {
require.Equal(t, "https://video.example/file?key=abc", ensureAPIKey("https://video.example/file", "abc"))
require.Equal(t, "https://video.example/file?alt=media&key=abc", ensureAPIKey("https://video.example/file?alt=media", "abc"))
require.Equal(t, "https://video.example/file?key=existing", ensureAPIKey("https://video.example/file?key=existing", "abc"))
require.Equal(t, "https://video.example/file", ensureAPIKey("https://video.example/file", ""))
}

+ 14
- 0
deploy/build_newapi.sh Просмотреть файл

@@ -0,0 +1,14 @@
#!/bin/bash
set -e
cd /mnt/d/code/new-api
BRANCH=$(git rev-parse --abbrev-ref HEAD | sed 's/\//-/g')
COMMIT=$(git rev-parse --short HEAD)
DIRTY=""
git diff --quiet && git diff --cached --quiet 2>/dev/null || DIRTY="-dirty"
TAG="$(date +%Y%m%d%H%M)-${BRANCH}-${COMMIT}${DIRTY}"
IMAGE="registry.cn-hangzhou.aliyuncs.com/fengsilin/new-api:${TAG}"
echo "Building: $IMAGE"
docker build -t "$IMAGE" . 2>&1 | tail -12
echo "===exit ${PIPESTATUS[0]}==="
echo "IMAGE=$IMAGE"
echo "TAG=$TAG"

+ 39
- 0
deploy/postgres-partman/Dockerfile Просмотреть файл

@@ -0,0 +1,39 @@
# Custom PostgreSQL 18.4 image with pg_partman + pg_cron extensions.
#
# The logs table of new-api is a RANGE-partitioned table on created_at, created
# and maintained entirely on the database side by the pg_partman extension.
# pg_cron drives periodic maintenance (premake future weekly partitions).
#
# Both extensions are installed via Debian packages. pg_cron must be loaded via
# shared_preload_libraries, so we append it to postgresql.conf.sample: initdb
# generates PGDATA/postgresql.conf from this sample, meaning both the temporary
# server (during docker-entrypoint-initdb.d) and the real server load pg_cron,
# allowing CREATE EXTENSION pg_cron to succeed at init time.
#
# Build/push:
# docker build -t registry.cn-hangzhou.aliyuncs.com/fengsilin/postgres-partman:18 .
# docker push registry.cn-hangzhou.aliyuncs.com/fengsilin/postgres-partman:18

FROM postgres:18

# Use Aliyun Debian mirror for faster/reliable apt downloads from within China.
# postgres:18 is based on Debian trixie which uses the deb822-format
# /etc/apt/sources.list.d/debian.sources file instead of sources.list.
RUN sed -i 's|deb.debian.org|mirrors.aliyun.com|g; s|security.debian.org|mirrors.aliyun.com|g' \
/etc/apt/sources.list.d/debian.sources

RUN apt-get update \
&& apt-get install -y --no-install-recommends \
postgresql-18-partman \
postgresql-18-cron \
&& rm -rf /var/lib/apt/lists/*

# Preload pg_cron so it is available during first-run init scripts.
# pg_partman is a plain SQL extension and does not need preloading.
RUN echo "shared_preload_libraries = 'pg_cron'" >> /usr/share/postgresql/18/postgresql.conf.sample
RUN echo "cron.database_name = 'new-api'" >> /usr/share/postgresql/18/postgresql.conf.sample

# First-run init: create extensions, the partitioned logs parent table, register
# it with pg_partman, build core indexes, and schedule maintenance via pg_cron.
# Runs against the database named by POSTGRES_DB (new-api).
COPY docker-entrypoint-initdb.d/01-partition.sql /docker-entrypoint-initdb.d/01-partition.sql

+ 102
- 0
deploy/postgres-partman/docker-entrypoint-initdb.d/01-partition.sql Просмотреть файл

@@ -0,0 +1,102 @@
-- 01-partition.sql
-- First-run initialization for new-api's partitioned logs table.
-- Executed by the postgres official entrypoint against the database named by
-- POSTGRES_DB (new-api), as a superuser.
--
-- This script is the SINGLE source of truth for the logs partitioning setup.
-- The new-api Go application does NOT create or maintain partitions; it only
-- skips GORM AutoMigrate for the logs table on PostgreSQL (see
-- model/main.go::ensureLogTable). pg_partman owns partition lifecycle, pg_cron
-- drives periodic maintenance.

-- 1. Extensions --------------------------------------------------------------
-- Install pg_partman into a dedicated schema named "partman" so its functions
-- are referenced as partman.xxx (pg_partman 5.x does not create the schema
-- automatically; without this CREATE EXTENSION lands in public and
-- partman.create_parent fails with "schema partman does not exist").
CREATE SCHEMA IF NOT EXISTS partman;
CREATE EXTENSION IF NOT EXISTS pg_partman WITH SCHEMA partman;
-- pg_cron must be created in the database set by cron.database_name (new-api),
-- which requires shared_preload_libraries='pg_cron' (configured in the image).
-- Include partman in the default search_path so its functions resolve without
-- schema-qualifying every call below.
SET search_path = partman, public;
CREATE EXTENSION IF NOT EXISTS pg_cron;

-- 2. Partitioned parent table -----------------------------------------------
-- Columns mirror model.Log exactly. PG requires the partition key (created_at)
-- to be part of the primary key, so it is a composite (id, created_at). The app
-- never queries logs by id alone, so this does not affect business logic.
CREATE TABLE IF NOT EXISTS public.logs (
id BIGSERIAL,
user_id INTEGER,
created_at BIGINT NOT NULL,
type INTEGER,
content TEXT,
username VARCHAR(64) DEFAULT '',
token_name VARCHAR(255) DEFAULT '',
model_name VARCHAR(255) DEFAULT '',
quota INTEGER DEFAULT 0,
prompt_tokens INTEGER DEFAULT 0,
completion_tokens INTEGER DEFAULT 0,
use_time INTEGER DEFAULT 0,
is_stream BOOLEAN,
channel_id INTEGER,
token_id INTEGER DEFAULT 0,
"group" VARCHAR(255),
ip VARCHAR(64) DEFAULT '',
request_id VARCHAR(64) DEFAULT '',
chat_id VARCHAR(128) DEFAULT '',
upstream_id VARCHAR(128) DEFAULT '',
other TEXT,
PRIMARY KEY (id, created_at)
) PARTITION BY RANGE (created_at);

-- 3. Hand off to pg_partman --------------------------------------------------
-- Weekly native range partitioning on the bigint epoch (seconds) column.
-- p_epoch='seconds' -> created_at is a unix-seconds bigint
-- p_type='range' -> pg_partman 5.x uses PG-native partitioning with
-- p_type values 'range'/'list' (the old 'native'
-- alias from 4.x was removed)
-- p_interval='1 week' -> one partition per week (pg_partman 5.x dropped
-- the 'weekly' preset in favor of native PG
-- interval values)
-- p_date_trunc_interval='week' -> align partition boundaries to ISO weeks (Monday)
-- p_premake=8 -> always keep 8 future weeks pre-created
-- p_default_table=true -> create a DEFAULT partition catching out-of-range
-- inserts so they never fail silently (RecordConsumeLog
-- only logs errors without retrying)
-- No retention is set: per current requirement we do NOT auto-drop old partitions.
SELECT partman.create_parent(
p_parent_table => 'public.logs',
p_control => 'created_at',
p_type => 'range',
p_interval => '1 week',
p_epoch => 'seconds',
p_date_trunc_interval => 'week',
p_premake => 8,
p_default_table => true
);

-- 4. Core indexes on the parent table ---------------------------------------
-- PG 11+ propagates indexes created on a partitioned parent to all child
-- partitions automatically. These cover the hot query paths in model/log.go
-- (GetAllLogs / GetUserLogs ordering by created_at desc, id desc; lookups by
-- user/model/channel/request/token). Trimmed from the model's 16 index tags to
-- the 6 actually used, cutting index write overhead.
CREATE INDEX IF NOT EXISTS idx_logs_created_at_id ON public.logs (created_at DESC, id DESC);
CREATE INDEX IF NOT EXISTS idx_logs_user_id_created ON public.logs (user_id, created_at DESC);
CREATE INDEX IF NOT EXISTS idx_logs_model_name ON public.logs (model_name);
CREATE INDEX IF NOT EXISTS idx_logs_channel_id ON public.logs (channel_id);
CREATE INDEX IF NOT EXISTS idx_logs_request_id ON public.logs (request_id);
CREATE INDEX IF NOT EXISTS idx_logs_token_id ON public.logs (token_id);

-- 5. Schedule periodic maintenance ------------------------------------------
-- run_maintenance_proc() inspects partman.part_config and premakes the next
-- partitions when needed. Every 30 minutes is more than enough; weekly partitions
-- only need creation roughly once a week. No retention => no drops.
SELECT cron.schedule(
'log-partition-maint',
'*/30 * * * *',
$$CALL partman.run_maintenance_proc()$$
);

+ 89
- 0
docs/superpowers/specs/2026-07-17-chinamobile-asset-isolation-design.md Просмотреть файл

@@ -0,0 +1,89 @@
# China Mobile Asset Library Per-User Isolation

## Goal

Give every platform user an isolated China Mobile AIGC asset group for each
resolved China Mobile channel. The platform creates and manages the group
internally. A user can only create, list, read, update, and delete assets in
that group.

## Scope

The change applies only after the asset proxy resolves a China Mobile asset
channel. Other asset adapters retain their existing action support and request
behavior.

Historical China Mobile groups and assets are not migrated. A user receives a
new group on their first asset request after rollout.

## Data Model

Add a persistent binding table with these fields:

- `user_id`: platform user ID.
- `channel_id`: resolved China Mobile channel ID.
- `group_id`: upstream China Mobile asset group ID.
- timestamps managed through the existing GORM conventions.

`(user_id, channel_id)` is unique. The group identity is deliberately scoped
to a channel because channels can represent different China Mobile credentials
or resource pools.

## Group Lifecycle

The asset proxy resolves the user and channel before dispatching a China Mobile
asset request. It calls a dedicated service to obtain the user group binding.

If no binding exists, the service creates an `AIGC` group through the internal
China Mobile adapter path, persists the returned group ID, and returns it. A
database uniqueness constraint handles concurrent first requests: after a
unique-conflict result, the service reads and returns the winning binding.

If upstream creation or binding persistence fails, the asset request fails. It
must never fall back to the shared upstream library. An orphaned upstream group
may remain when persistence fails after creation; it is inaccessible through
the platform and does not weaken isolation.

## Request Authorization

Clients are forbidden from calling all group actions:

- `CreateAssetGroup`
- `ListAssetGroups`
- `GetAssetGroup`
- `UpdateAssetGroup`
- `DeleteAssetGroup`

The proxy returns HTTP 403 before dispatching these actions for a resolved
China Mobile channel. Internal group creation bypasses this public-action gate.

For permitted asset actions, the proxy injects the bound group ID and does not
trust user-supplied group restrictions:

- `CreateAsset`: overwrite `GroupId`.
- `ListAssets`: overwrite `Filter.GroupIds` with the bound group ID.
- `GetAsset`, `UpdateAsset`, and `DeleteAsset`: retrieve the asset first and
verify its upstream `GroupId` equals the bound group ID. Treat a mismatch or
missing asset as not found, without exposing another user's asset details.

## Error Behavior

Group provisioning failures and binding-store failures stop the request and
return an explicit service or upstream error. Authorization failures for group
actions return 403. Cross-user asset IDs return the same not-found outcome as
an absent asset ID.

## Testing

Add focused controller and service tests for:

- initial group creation and reuse for the same user and channel;
- separate bindings for distinct users or channels;
- concurrent initial requests converging on one persisted binding;
- client-supplied `GroupId` and `Filter.GroupIds` being overwritten;
- cross-user get, update, and delete attempts returning not found;
- all five public group actions returning 403 only for China Mobile;
- regression coverage confirming non-China-Mobile adapters preserve current
group-action behavior.

Run focused controller/service tests, then the relevant package test suite.

+ 329
- 0
docs/testing/2026-07-21-cn-tianyiyun-seedance-mini-e2e.md Просмотреть файл

@@ -0,0 +1,329 @@
# 国内站天翼云 Seedance Mini 在线测试记录

测试时间:2026-07-21(Asia/Shanghai)

测试环境:国内站 `https://router.lanqi.tech`,部署镜像
`registry.cn-hangzhou.aliyuncs.com/fengsilin/new-api:202607211136-master-fedd394e-dirty`。

## 脱敏说明

- 所有请求均使用用户提供的 Bearer Token;本文不记录 Token。
- 天翼云素材临时下载 URL 中的查询参数、凭证、签名和安全令牌均已移除。
- 测试素材 ID 与任务 ID 保留,便于在国内站数据库和天翼云控制台审计。

## 前置核对

国内站实际数据源中存在启用的天翼云 type 60 渠道:

| 字段 | 值 |
| --- | --- |
| 渠道 ID | 16 |
| 状态 | 启用 |
| 模型 | `Doubao-Seedance-2.0`、`Doubao-Seedance-2.0-fast`、`Doubao-Seedance-2.0-mini` |
| 用户素材绑定数 | 2 |

结论:本次请求确实可路由到国内站天翼云渠道。

## 1. CreateAsset:公网图片上传

请求:

```http
POST /api/v1/volcengine/asset?Action=CreateAsset HTTP/1.1
Host: router.lanqi.tech
Authorization: Bearer <redacted>
Content-Type: application/json

{"URL":"https://ark-project.tos-cn-beijing.volces.com/doc_image/r2v_tea_pic1.jpg","AssetType":"Image","Name":"new-api-tianyiyun-e2e-20260721"}
```

响应(HTTP 200):

```json
{
"ResponseMetadata": {
"RequestId": "b63e4b6b-e65e-46b1-96e2-cb6247382e21",
"Action": "CreateAsset",
"Version": "2024-01-01",
"Service": "ark",
"Region": "cn-beijing"
},
"Result": "asset-20260721141127-zdl8v"
}
```

结果:通过。网关成功将既有 Action 接口映射到天翼云素材上传接口,并返回天翼云素材 ID。

## 2. GetAsset:素材查询

请求:

```http
POST /api/v1/volcengine/asset?Action=GetAsset HTTP/1.1
Host: router.lanqi.tech
Authorization: Bearer <redacted>
Content-Type: application/json

{"Id":"asset-20260721141127-zdl8v"}
```

响应(HTTP 200,临时 URL 已脱敏):

```json
{
"ResponseMetadata": {
"RequestId": "86327f75-fd21-40a5-a78e-e200d7111ba2",
"Action": "GetAsset",
"Version": "2024-01-01",
"Service": "ark",
"Region": "cn-beijing"
},
"Result": {
"AssetType": "Image",
"CreatedAt": "2026-07-21 14:11:27",
"GroupId": "group-20260720122310-q4lcv",
"Id": "asset-20260721141127-zdl8v",
"Name": "new-api-tianyiyun-e2e-20260721",
"Status": "Active",
"URL": "https://ark-media-asset.tos-cn-beijing.volces.com/<redacted>?<signed-query-redacted>",
"UpdatedAt": "2026-07-21 14:11:28"
}
}
```

结果:通过。素材状态为 `Active`,可作为后续视频任务的 `asset://asset-20260721141127-zdl8v` 输入。

## 3. CreateAsset:data URI 本地拒绝

请求:

```http
POST /api/v1/volcengine/asset?Action=CreateAsset HTTP/1.1
Host: router.lanqi.tech
Authorization: Bearer <redacted>
Content-Type: application/json

{"URL":"data:image/png;base64,AAAA","AssetType":"Image","Name":"should-not-upload"}
```

响应(HTTP 400):

```json
{
"error": {
"message": "asset URL must use http or https",
"type": "invalid_request_error"
}
}
```

结果:通过。请求在网关本地拒绝,没有创建天翼云素材。

## 4. Mini 视频任务:OpenAI 视频入口兼容性

请求:

```http
POST /v1/video/generations HTTP/1.1
Host: router.lanqi.tech
Authorization: Bearer <redacted>
Content-Type: application/json

{
"model": "Doubao-Seedance-2.0-mini",
"content": [
{"type": "text", "text": "A five-second product showcase of a red apple on a clean white background. No text overlay."},
{"type": "image_url", "role": "reference_image", "image_url": {"url": "asset://asset-20260721141127-zdl8v"}}
],
"ratio": "16:9",
"duration": 5,
"generate_audio": false,
"watermark": false
}
```

响应(HTTP 400):

```json
{
"code": "invalid_request",
"message": "prompt is required",
"data": null
}
```

结果:该入口要求顶层 `prompt`,不接受此天翼云原生 `content` 格式;未创建上游任务、未产生推理费用。

## 5. Mini 视频任务:天翼云原生兼容入口

请求:

```http
POST /api/v3/contents/generations/tasks HTTP/1.1
Host: router.lanqi.tech
Authorization: Bearer <redacted>
Content-Type: application/json

{
"model": "Doubao-Seedance-2.0-mini",
"content": [
{"type": "text", "text": "A five-second product showcase of a red apple on a clean white background. No text overlay."},
{"type": "image_url", "role": "reference_image", "image_url": {"url": "asset://asset-20260721141127-zdl8v"}}
],
"ratio": "16:9",
"duration": 5,
"generate_audio": false,
"watermark": false
}
```

响应(HTTP 400):

```json
{
"code": "pricing_no_match",
"message": "pricing dimensions did not match any row",
"data": null
}
```

结果:Matrix usage billing 模型能力校验已经通过(未出现此前的
`model ... does not support matrix usage billing`),但该公开模型目前没有命中线上
Matrix 定价表的维度行。请求在调用天翼云之前被拒绝,因此没有视频任务、上游任务 ID
或 token 用量可供轮询和结算验证,也没有产生 Mini 推理费用。

## 6. Mini 视频任务:补充分辨率后重试

第 5 节请求遗漏了 `resolution`。天翼云 Mini 仅支持 480P 和 720P,本次补充
`resolution: "480p"` 后使用相同素材重试:

```http
POST /api/v3/contents/generations/tasks HTTP/1.1
Host: router.lanqi.tech
Authorization: Bearer <redacted>
Content-Type: application/json

{
"model": "Doubao-Seedance-2.0-mini",
"content": [
{"type": "text", "text": "A five-second product showcase of a red apple on a clean white background. No text overlay."},
{"type": "image_url", "role": "reference_image", "image_url": {"url": "asset://asset-20260721141127-zdl8v"}}
],
"resolution": "480p",
"ratio": "16:9",
"duration": 5,
"generate_audio": false,
"watermark": false
}
```

响应(HTTP 401):

```json
{
"code": "fail_to_fetch_task",
"message": "{\n \"code\": 401,\n \"message\": \"internal-error\",\n \"detail\": \"api-key access forbit\",\n \"error\": {\n \"message\": \"api-key access forbit\",\n \"type\": \"internal-error\",\n \"code\": \"401\"\n }\n}\n",
"data": null
}
```

结果:补充分辨率后不再返回 `pricing_no_match`,证明 Matrix 定价行已命中、公开
Mini 模型到天翼云上游模型的映射也已通过。请求已实际到达天翼云,但渠道配置的上游
API Key 没有该模型调用权限(`api-key access forbit`)。没有创建成功视频任务,因而
没有可轮询的 `usage.total_tokens` 或结算记录。

## 后续动作

1. 在天翼云客户控制台确认该渠道 API Key 已开通并加白 `cdance2.0-mini-0611`;必要时按天翼云流程申请 Seedance 2.0 Mini 权限。
2. 权限开通后,保持第 6 节的请求体不变重试;成功后记录公共任务 ID、轮询响应中的 `usage.total_tokens`、预扣日志、settlement 日志、用户额度变化和 `tasks.quota`。
3. 继续保留本次素材 `asset-20260721141127-zdl8v`,可直接用于重试,避免重复上传。

## 7. Fast 模型权限检查

使用第 6 节相同的素材、`resolution: "480p"`、`ratio: "16:9"`、`duration: 5`,仅将模型替换为
`Doubao-Seedance-2.0-fast`。

响应(HTTP 401):

```json
{
"code": "fail_to_fetch_task",
"message": "{\n \"code\": 401,\n \"message\": \"internal-error\",\n \"detail\": \"api-key access forbit\"\n}",
"data": null
}
```

结果:Fast 与 Mini 都被同一个天翼云渠道上游 API Key 拒绝。问题是上游模型权限,非网关模型映射或 Matrix 定价问题。

## 8. 标准 Seedance 2.0 完整提交、轮询与结算

使用以下请求验证该渠道的标准模型权限和 Matrix 结算链路:

```http
POST /api/v3/contents/generations/tasks HTTP/1.1
Host: router.lanqi.tech
Authorization: Bearer <redacted>
Content-Type: application/json

{
"model": "Doubao-Seedance-2.0",
"content": [
{"type": "text", "text": "A five-second product showcase of a red apple on a clean white background. No text overlay."},
{"type": "image_url", "role": "reference_image", "image_url": {"url": "asset://asset-20260721141127-zdl8v"}}
],
"resolution": "480p",
"ratio": "16:9",
"duration": 5,
"generate_audio": false,
"watermark": false
}
```

提交响应(HTTP 200):

```json
{
"created_at": 1784615279,
"id": "task_UwH2BAtzuR8KI4Xs6KIsN3TxS9k77eQv",
"model": "Doubao-Seedance-2.0"
}
```

轮询:前 6 次为 `running`;第 7 次成功。成功响应(临时下载 URL 已脱敏):

```json
{
"content": {
"video_url": "https://ark-acg-cn-beijing.tos-cn-beijing.volces.com/<redacted>?<signed-query-redacted>"
},
"created_at": 1784615279,
"duration": 5,
"framespersecond": 24,
"generate_audio": false,
"id": "task_UwH2BAtzuR8KI4Xs6KIsN3TxS9k77eQv",
"model": "doubao-seedance-2-0-260128",
"ratio": "16:9",
"resolution": "480p",
"status": "succeeded",
"updated_at": 1784615379,
"usage": {
"completion_tokens": 50638,
"total_tokens": 50638
}
}
```

数据库结算核对:

| 项目 | 值 |
| --- | ---: |
| 任务状态 | `SUCCESS` |
| 任务最终 quota | 166382 |
| 预扣 quota | 1642 |
| 结算补扣 quota | 164740 |
| 实际总 quota | 166382 |
| completion / total tokens | 50638 / 50638 |
| Matrix 单价 | 6.571428571429 USD / 1M tokens |
| 分组倍率 | 1 |

结果:标准模型权限正常;素材 `asset://` 输入、天翼云异步轮询、usage token 解析、Matrix 差额结算、消费日志和 `tasks.quota` 持久化全部验证通过。

+ 30
- 0
docs/testing/2026-07-24-admin-video-channel-binding.md Просмотреть файл

@@ -0,0 +1,30 @@
# 管理员视频渠道绑定验收记录

日期:2026-07-24

## 自动化验证

- `go test ./controller -run '^TestAdmin(Set|Get)UserVideoChannelBindings' -count=1`:通过。
- GET 只返回用户的具体 Token 分组,不含 `auto`,且不泄露渠道密钥。
- PUT 可切换 Seedance 渠道;省略的家族绑定被清空。
- 拒绝非用户分组、`auto`、错误家族、跨分组、禁用渠道和重复绑定;失败不改变原绑定。
- 仅覆盖当前 Token 分组,历史 `legacy` 分组的绑定保持不变。
- `go test ./...`:通过。
- `cd web; bun run build`:通过。保留项目既有的循环分包与大包告警。

## 手工验收步骤

1. 用管理员身份打开“用户管理”,编辑一个同时拥有 `default`、`vip` 和 `auto` Token 的用户。
2. 在“视频渠道绑定”卡片中点击“配置视频渠道”。确认只出现 `default`、`vip`,每个分组显示 Seedance 和 Kling。
3. 为 `default / Seedance` 选择一个候选渠道并保存;重新打开弹窗,确认选择仍在。
4. 清空 `default / Seedance` 并保存;重新打开弹窗,确认显示“未绑定”。
5. 使用该用户 `default` Token 发起对应 Seedance 请求:有绑定时命中指定渠道;清空后恢复既有自动选择。
6. 准备一个无任何 Seedance/Kling 候选渠道的 Token group,确认整个 group 不显示。
7. 准备一个只有 Seedance 候选渠道的 Token group,确认显示该 group 和 Seedance,但不显示 Kling。
8. 当所有具体 Token group 均无候选渠道时,确认弹窗显示“该用户当前 Token 分组没有可用的视频渠道”。

记录请求时只保存用户 ID、Token group、模型、渠道 ID 和 HTTP 状态,禁止保存 Token、渠道 Key、AK/SK 或签名 URL。

## 已知验证限制

`bun run i18n:extract`、`bun run i18n:sync` 与 `bun run i18n:lint` 当前均因依赖版本不匹配失败:`react-i18next` 请求 `i18next.keyFromSelector`,但安装的 `i18next` 未导出该符号。新增中英文键已手动同步,生产构建已验证通过。

+ 13
- 1
go.mod Просмотреть файл

@@ -50,6 +50,8 @@ require (
github.com/tidwall/sjson v1.2.5
github.com/tiktoken-go/tokenizer v0.6.2
github.com/yapingcat/gomedia v0.0.0-20240906162731-17feea57090c
gitlab.ecloud.com/ecloud/ecloudsdkcore v1.0.6
gitlab.ecloud.com/ecloud/ecloudsdkmaas v1.0.2
golang.org/x/crypto v0.48.0
golang.org/x/image v0.23.0
golang.org/x/net v0.49.0
@@ -60,6 +62,7 @@ require (
gorm.io/driver/mysql v1.4.3
gorm.io/driver/postgres v1.5.2
gorm.io/gorm v1.25.2
maas_seedance_sdk_1.0.0_go v1.0.0
)

require (
@@ -110,6 +113,7 @@ require (
github.com/jfreymuth/vorbis v1.0.2 // indirect
github.com/jinzhu/inflection v1.0.0 // indirect
github.com/jinzhu/now v1.1.5 // indirect
github.com/jmespath/go-jmespath v0.4.0 // indirect
github.com/json-iterator/go v1.1.12 // indirect
github.com/klauspost/compress v1.18.0 // indirect
github.com/klauspost/cpuid/v2 v2.3.0 // indirect
@@ -123,6 +127,7 @@ require (
github.com/modern-go/reflect2 v1.0.2 // indirect
github.com/munnerz/goautoneg v0.0.0-20191010083416-a7dc8b61c822 // indirect
github.com/ncruces/go-strftime v0.1.9 // indirect
github.com/openai/openai-go/v3 v3.1.0 // indirect
github.com/pelletier/go-toml/v2 v2.2.1 // indirect
github.com/pmezard/go-difflib v1.0.1-0.20181226105442-5d4384ee4fb2 // indirect
github.com/prometheus/client_model v0.6.1 // indirect
@@ -131,18 +136,25 @@ require (
github.com/remyoudompheng/bigfft v0.0.0-20230129092748-24d4a6f8daec // indirect
github.com/samber/go-singleflightx v0.3.2 // indirect
github.com/tidwall/match v1.1.1 // indirect
github.com/tidwall/pretty v1.2.0 // indirect
github.com/tidwall/pretty v1.2.1 // indirect
github.com/tklauser/go-sysconf v0.3.12 // indirect
github.com/tklauser/numcpus v0.6.1 // indirect
github.com/twitchyliquid64/golang-asm v0.15.1 // indirect
github.com/ugorji/go/codec v1.2.12 // indirect
github.com/volcengine/volc-sdk-golang v1.0.23 // indirect
github.com/volcengine/volcengine-go-sdk v1.2.10 // indirect
github.com/x448/float16 v0.8.4 // indirect
github.com/yusufpapurcu/wmi v1.2.3 // indirect
golang.org/x/arch v0.21.0 // indirect
golang.org/x/exp v0.0.0-20250620022241-b7579e27df2b // indirect
google.golang.org/protobuf v1.36.5 // indirect
gopkg.in/yaml.v2 v2.4.0 // indirect
modernc.org/libc v1.66.10 // indirect
modernc.org/mathutil v1.7.1 // indirect
modernc.org/memory v1.11.0 // indirect
modernc.org/sqlite v1.40.1 // indirect
)

replace maas_seedance_sdk_1.0.0_go => ./third_party/maas_seedance_sdk_1.0.0_go
replace gitlab.ecloud.com/ecloud/ecloudsdkcore => ./third_party/ecloudsdkcore
replace gitlab.ecloud.com/ecloud/ecloudsdkmaas => ./third_party/ecloudsdkmaas

+ 80
- 1
go.sum Просмотреть файл

@@ -1,3 +1,5 @@
cloud.google.com/go v0.26.0/go.mod h1:aQUYkXzVsufM+DwF1aE+0xfcU+56JwCaLick0ClmMTw=
github.com/BurntSushi/toml v0.3.1/go.mod h1:xHWCNGjB5oqiDr8zfno3MHue2Ht5sIBksp03qcyfWMU=
github.com/BurntSushi/toml v1.6.0 h1:dRaEfpa2VI55EwlIW72hMRHdWouJeRF7TPYhI+AUQjk=
github.com/BurntSushi/toml v1.6.0/go.mod h1:ukJfTF/6rtPPRCnwkur4qwRxa8vTRFBF0uk2lLoLwho=
github.com/Calcium-Ion/go-epay v0.0.4 h1:C96M7WfRLadcIVscWzwLiYs8etI1wrDmtFMuK2zP22A=
@@ -12,6 +14,7 @@ github.com/anknown/ahocorasick v0.0.0-20190904063843-d75dbd5169c0 h1:onfun1RA+Kc
github.com/anknown/ahocorasick v0.0.0-20190904063843-d75dbd5169c0/go.mod h1:4yg+jNTYlDEzBjhGS96v+zjyA3lfXlFd5CiTLIkPBLI=
github.com/anknown/darts v0.0.0-20151216065714-83ff685239e6 h1:HblK3eJHq54yET63qPCTJnks3loDse5xRmmqHgHzwoI=
github.com/anknown/darts v0.0.0-20151216065714-83ff685239e6/go.mod h1:pbiaLIeYLUbgMY1kwEAdwO6UKD5ZNwdPGQlwokS9fe8=
github.com/avast/retry-go v3.0.0+incompatible/go.mod h1:XtSnn+n/sHqQIpZ10K1qAevBhOOCWBLXXy3hyiqqBrY=
github.com/aws/aws-sdk-go-v2 v1.37.2 h1:xkW1iMYawzcmYFYEV0UCMxc8gSsjCGEhBXQkdQywVbo=
github.com/aws/aws-sdk-go-v2 v1.37.2/go.mod h1:9Q0OoGQoboYIAJyslFyF1f5K1Ryddop8gqMhWx/n4Wg=
github.com/aws/aws-sdk-go-v2/aws/protocol/eventstream v1.7.0 h1:6GMWV6CNpA/6fbFHnoAjrv4+LGfyTqZz2LtCHnspgDg=
@@ -37,8 +40,10 @@ github.com/bytedance/sonic v1.14.1 h1:FBMC0zVz5XUmE4z9wF4Jey0An5FueFvOsTKKKtwIl7
github.com/bytedance/sonic v1.14.1/go.mod h1:gi6uhQLMbTdeP0muCnrjHLeCUPyb70ujhnNlhOylAFc=
github.com/bytedance/sonic/loader v0.3.0 h1:dskwH8edlzNMctoruo8FPTJDF3vLtDT0sXZwvZJyqeA=
github.com/bytedance/sonic/loader v0.3.0/go.mod h1:N8A3vUdtUebEY2/VQC0MyhYeKUFosQU6FxH2JmUe6VI=
github.com/census-instrumentation/opencensus-proto v0.2.1/go.mod h1:f6KPmirojxKA12rnyqOA5BBL4O983OfeGPqjHWSTneU=
github.com/cespare/xxhash/v2 v2.3.0 h1:UL815xU9SqsFlibzuggzjXhog7bL6oX9BbNZnL2UFvs=
github.com/cespare/xxhash/v2 v2.3.0/go.mod h1:VGX0DQ3Q6kWi7AoAeZDth3/j3BFtOZR5XLFGgcrjCOs=
github.com/client9/misspell v0.3.4/go.mod h1:qj6jICC3Q7zFZvVWo7KLAzC3yx5G7kyvSDkc90ppPyw=
github.com/cloudwego/base64x v0.1.6 h1:t11wG9AECkCDk5fMSoxmufanudBtJ+/HemLstXDLI2M=
github.com/cloudwego/base64x v0.1.6/go.mod h1:OFcloc187FXDaYHvrNIjxSe8ncn0OOM8gEHfghB2IPU=
github.com/creack/pty v1.1.7/go.mod h1:lj5s0c3V2DBrqTV7llrYr5NG6My20zk30Fl46Y7DoTY=
@@ -53,6 +58,8 @@ github.com/dlclark/regexp2 v1.11.5 h1:Q/sSnsKerHeCkc/jSTNq1oCm7KiVgUMZRDUoRu0JQZ
github.com/dlclark/regexp2 v1.11.5/go.mod h1:DHkYz0B9wPfa6wondMfaivmHpzrQ3v9q8cnmRbL6yW8=
github.com/dustin/go-humanize v1.0.1 h1:GzkhY7T5VNhEkwH0PVJgjz+fX1rhBrR7pRT3mDkpeCY=
github.com/dustin/go-humanize v1.0.1/go.mod h1:Mu1zIs6XwVuF/gI1OepvI0qD18qycQx+mFykh5fBlto=
github.com/envoyproxy/go-control-plane v0.9.1-0.20191026205805-5f8ba28d4473/go.mod h1:YTl/9mNaCwkRvm6d1a2C3ymFceY/DCBVvsKhRF0iEA4=
github.com/envoyproxy/protoc-gen-validate v0.1.0/go.mod h1:iSmxcyjqTsJpI2R4NaDN7+kN2VEUnK/pcBlmesArF7c=
github.com/fsnotify/fsnotify v1.4.9 h1:hsms1Qyu0jgnwNXIxa+/V/PDsU6CfLf6CNO8H7IWoS4=
github.com/fsnotify/fsnotify v1.4.9/go.mod h1:znqG4EE+3YCdAaPaxE2ZRY/06pZUdp0tY4IgpuI1SZQ=
github.com/fxamacker/cbor/v2 v2.9.0 h1:NpKPmjDBgUfBms6tr6JZkTHtfFGcMKsw3eGcmD/sapM=
@@ -133,10 +140,26 @@ github.com/golang-jwt/jwt/v5 v5.3.0 h1:pv4AsKCKKZuqlgs5sUmn4x8UlGa0kEVt/puTpKx9v
github.com/golang-jwt/jwt/v5 v5.3.0/go.mod h1:fxCRLWMO43lRc8nhHWY6LGqRcf+1gQWArsqaEUEa5bE=
github.com/golang/freetype v0.0.0-20170609003504-e2365dfdc4a0 h1:DACJavvAHhabrF08vX0COfcOBJRhZ8lUbR+ZWIs0Y5g=
github.com/golang/freetype v0.0.0-20170609003504-e2365dfdc4a0/go.mod h1:E/TSTwGwJL78qG/PmXZO1EjYhfJinVAhrmmHX6Z8B9k=
github.com/golang/glog v0.0.0-20160126235308-23def4e6c14b/go.mod h1:SBH7ygxi8pfUlaOkMMuAQtPIUF8ecWP5IEl/CR7VP2Q=
github.com/golang/groupcache v0.0.0-20210331224755-41bb18bfe9da h1:oI5xCqsCo564l8iNU+DwB5epxmsaqB+rhGL0m5jtYqE=
github.com/golang/groupcache v0.0.0-20210331224755-41bb18bfe9da/go.mod h1:cIg4eruTrX1D+g88fzRXU5OdNfaM+9IcxsU14FzY7Hc=
github.com/golang/mock v1.1.1/go.mod h1:oTYuIxOrZwtPieC+H1uAHpcLFnEyAGVDL/k47Jfbm0A=
github.com/golang/protobuf v1.2.0/go.mod h1:6lQm79b+lXiMfvg/cZm0SGofjICqVBUtrP5yJMmIC1U=
github.com/golang/protobuf v1.3.2/go.mod h1:6lQm79b+lXiMfvg/cZm0SGofjICqVBUtrP5yJMmIC1U=
github.com/golang/protobuf v1.3.3/go.mod h1:vzj43D7+SQXF/4pzW/hwtAqwc6iTitCiVSaWz5lYuqw=
github.com/golang/protobuf v1.4.0-rc.1/go.mod h1:ceaxUfeHdC40wWswd/P6IGgMaK3YpKi5j83Wpe3EHw8=
github.com/golang/protobuf v1.4.0-rc.1.0.20200221234624-67d41d38c208/go.mod h1:xKAWHe0F5eneWXFV3EuXVDTCmh+JuBKY0li0aMyXATA=
github.com/golang/protobuf v1.4.0-rc.2/go.mod h1:LlEzMj4AhA7rCAGe4KMBDvJI+AwstrUpVNzEA03Pprs=
github.com/golang/protobuf v1.4.0-rc.4.0.20200313231945-b860323f09d0/go.mod h1:WU3c8KckQ9AFe+yFwt9sWVRKCVIyN9cPHBJSNnbL67w=
github.com/golang/protobuf v1.4.0/go.mod h1:jodUvKwWbYaEsadDk5Fwe5c77LiNKVO9IDvqG2KuDX0=
github.com/golang/protobuf v1.4.1/go.mod h1:U8fpvMrcmy5pZrNK1lt4xCsGvpyWQ/VVv6QDs8UjoX8=
github.com/golang/protobuf v1.4.3/go.mod h1:oDoupMAO8OvCJWAcko0GGGIgR6R6ocIYbsSw735rRwI=
github.com/golang/protobuf v1.5.0/go.mod h1:FsONVRAS9T7sI+LIUmWTfcYkHO4aIWwzhcaSAoJOfIk=
github.com/google/go-cmp v0.2.0/go.mod h1:oXzfMopK8JAjlY9xF4vHSVASa0yLyX7SntLO5aqRK0M=
github.com/google/go-cmp v0.3.0/go.mod h1:8QqcDgzrUqlUb/G2PQTWiueGozuR1884gddMywk6iLU=
github.com/google/go-cmp v0.3.1/go.mod h1:8QqcDgzrUqlUb/G2PQTWiueGozuR1884gddMywk6iLU=
github.com/google/go-cmp v0.4.0/go.mod h1:v8dTdLbMG2kIc/vJvl+f65V22dbkXbowE6jgT/gNBxE=
github.com/google/go-cmp v0.5.0/go.mod h1:v8dTdLbMG2kIc/vJvl+f65V22dbkXbowE6jgT/gNBxE=
github.com/google/go-cmp v0.5.5/go.mod h1:v8dTdLbMG2kIc/vJvl+f65V22dbkXbowE6jgT/gNBxE=
github.com/google/go-cmp v0.6.0/go.mod h1:17dUlkBOakJ0+DkrSSNjCkIjxS6bF9zb3elmeNGIjoY=
github.com/google/go-cmp v0.7.0 h1:wk8382ETsv4JYUZwIsn6YpYiWiBsYLSJiTsyBybVuN8=
@@ -147,6 +170,7 @@ github.com/google/gofuzz v1.0.0/go.mod h1:dBl0BpW6vV/+mYPU4Po3pmUjxk6FQPldtuIdl/
github.com/google/pprof v0.0.0-20250317173921-a4b03ec1a45e h1:ijClszYn+mADRFY17kjQEVQ1XRhq2/JR1M3sGqeJoxs=
github.com/google/pprof v0.0.0-20250317173921-a4b03ec1a45e/go.mod h1:boTsfXsheKC2y+lKOCMpSfarhxDeIzfZG1jqGcPl3cA=
github.com/google/uuid v1.1.2/go.mod h1:TIyPZe4MgqvfeYDBFedMoGGpEw/LqOeaOT+nhxU+yHo=
github.com/google/uuid v1.3.0/go.mod h1:TIyPZe4MgqvfeYDBFedMoGGpEw/LqOeaOT+nhxU+yHo=
github.com/google/uuid v1.6.0 h1:NIvaJDMOsjHA8n1jAhLSgzrAzy1Hgr+hNrb57e+94F0=
github.com/google/uuid v1.6.0/go.mod h1:TIyPZe4MgqvfeYDBFedMoGGpEw/LqOeaOT+nhxU+yHo=
github.com/gorilla/context v1.1.1 h1:AWwleXJkX/nhcU9bZSnZoi3h/qGYqQAGhq6zZe/aQW8=
@@ -184,6 +208,10 @@ github.com/jinzhu/inflection v1.0.0/go.mod h1:h+uFLlag+Qp1Va5pdKtLDYj+kHp5pxUVkr
github.com/jinzhu/now v1.1.4/go.mod h1:d3SSVoowX0Lcu0IBviAWJpolVfI5UJVZZ7cO71lE/z8=
github.com/jinzhu/now v1.1.5 h1:/o9tlHleP7gOFmsnYNz3RGnqzefHA47wQpKrrdTIwXQ=
github.com/jinzhu/now v1.1.5/go.mod h1:d3SSVoowX0Lcu0IBviAWJpolVfI5UJVZZ7cO71lE/z8=
github.com/jmespath/go-jmespath v0.4.0 h1:BEgLn5cpjn8UN1mAw4NjwDrS35OdebyEtFe+9YPoQUg=
github.com/jmespath/go-jmespath v0.4.0/go.mod h1:T8mJZnbsbmF+m6zOOFylbeCJqk5+pHWvzYPziyZiYoo=
github.com/jmespath/go-jmespath/internal/testify v1.5.1 h1:shLQSRRSCCPj3f2gpwzGwWFoC7ycTf1rcQZHOlsJ6N8=
github.com/jmespath/go-jmespath/internal/testify v1.5.1/go.mod h1:L3OGu8Wl2/fWfCI6z80xFu9LTZmf1ZRjMHUOPmWr69U=
github.com/joho/godotenv v1.5.1 h1:7eLL/+HRGLY0ldzfGMeQkb7vMd0as4CfYvUVzLqw0N0=
github.com/joho/godotenv v1.5.1/go.mod h1:f4LDr5Voq0i2e/R5DDNOoa2zzDfwtkZa6DnEwAbqwq4=
github.com/json-iterator/go v1.1.9/go.mod h1:KdQUCv79m/52Kvf8AW2vK1V8akMuk1QjK/uOdHXbAo4=
@@ -194,6 +222,7 @@ github.com/klauspost/compress v1.18.0/go.mod h1:2Pp+KzxcywXVXMr50+X0Q/Lsb43OQHYW
github.com/klauspost/cpuid/v2 v2.3.0 h1:S4CRMLnYUhGeDFDqkGriYKdfoFlDnMtqTiI/sFzhA9Y=
github.com/klauspost/cpuid/v2 v2.3.0/go.mod h1:hqwkgyIinND0mEev00jJYCxPNVRVXFQeu1XKlok6oO0=
github.com/kr/pretty v0.1.0/go.mod h1:dAy3ld7l9f0ibDNOQOHHMYYIIbhfbHSm3C4ZsoJORNo=
github.com/kr/pretty v0.2.0/go.mod h1:ipq/a2n7PKx3OHsz4KJII5eveXtPO4qwEXGdVfWzfnI=
github.com/kr/pretty v0.2.1/go.mod h1:ipq/a2n7PKx3OHsz4KJII5eveXtPO4qwEXGdVfWzfnI=
github.com/kr/pretty v0.3.0/go.mod h1:640gp4NfQd8pI5XOwp5fnNeVWj67G7CFk/SaSQn7NBk=
github.com/kr/pretty v0.3.1 h1:flRD4NNwYAUpkphVc1HcthR4KEIFJ65n8Mw5qdRn3LE=
@@ -242,6 +271,8 @@ github.com/onsi/ginkgo v1.16.5 h1:8xi0RTUf59SOSfEtZMvwTvXYMzG4gV23XVHOZiXNtnE=
github.com/onsi/ginkgo v1.16.5/go.mod h1:+E8gABHa3K6zRBolWtd+ROzc/U5bkGt0FwiG042wbpU=
github.com/onsi/gomega v1.18.1 h1:M1GfJqGRrBrrGGsbxzV5dqM2U2ApXefZCQpkukxYRLE=
github.com/onsi/gomega v1.18.1/go.mod h1:0q+aL8jAiMXy9hbwj2mr5GziHiwhAIQpFmmtT5hitRs=
github.com/openai/openai-go/v3 v3.1.0 h1:sBf6OYL6Pj1qMAkQEmkz8r8z+EBes+iI7gCuCgr8e/A=
github.com/openai/openai-go/v3 v3.1.0/go.mod h1:UOpNxkqC9OdNXNUfpNByKOtB4jAL0EssQXq5p8gO0Xs=
github.com/orcaman/writerseeker v0.0.0-20200621085525-1d3f536ff85e h1:s2RNOM/IGdY0Y6qfTeUKhDawdHDpK9RGBdx80qN4Ttw=
github.com/orcaman/writerseeker v0.0.0-20200621085525-1d3f536ff85e/go.mod h1:nBdnFKj15wFbf94Rwfq4m30eAcyY9V/IyKAGQFtqkW0=
github.com/pelletier/go-toml/v2 v2.0.1/go.mod h1:r9LEWfGN8R5k0VXJ+0BkIe7MYkRdwZOjgMj2KwnJFUo=
@@ -257,6 +288,7 @@ github.com/pquerna/otp v1.5.0 h1:NMMR+WrmaqXU4EzdGJEE1aUUI0AMRzsp96fFFWNPwxs=
github.com/pquerna/otp v1.5.0/go.mod h1:dkJfzwRKNiegxyNb54X/3fLwhCynbMspSyWKnvi1AEg=
github.com/prometheus/client_golang v1.22.0 h1:rb93p9lokFEsctTys46VnV1kLCDpVZ0a/Y92Vm0Zc6Q=
github.com/prometheus/client_golang v1.22.0/go.mod h1:R7ljNsLXhuQXYZYtw6GAE9AZg8Y7vEW5scdCXrWRXC0=
github.com/prometheus/client_model v0.0.0-20190812154241-14fe0d1b01d4/go.mod h1:xMI15A0UPsDsEKsMN9yxemIoYk6Tm2C1GtYGdfGttqA=
github.com/prometheus/client_model v0.6.1 h1:ZKSh/rekM+n3CeS952MLRAdFwIKqeY8b62p8ais2e9E=
github.com/prometheus/client_model v0.6.1/go.mod h1:OrxVMOVHjw3lKMa8+x6HeMGkHMQyHDk9E3jmP2AmGiY=
github.com/prometheus/common v0.62.0 h1:xasJaQlnWAeyHdUBeGjXmutelfJHWMRr+Fg4QszZ2Io=
@@ -286,6 +318,7 @@ github.com/stretchr/objx v0.5.2 h1:xuMeJ0Sdp5ZMRXx/aWO6RZxdr3beISkG5/G/aIRr3pY=
github.com/stretchr/objx v0.5.2/go.mod h1:FRsXN1f5AsAjCGJKqEizvkpNtU+EGNCLh3NxZ/8L+MA=
github.com/stretchr/testify v1.3.0/go.mod h1:M5WIy9Dh21IEIfnGCwXGc5bZfKNJtfHm1UVUgZn+9EI=
github.com/stretchr/testify v1.4.0/go.mod h1:j7eGeouHqKxXV5pUuKE4zz7dFj8WfuZ+81PSLYec5m4=
github.com/stretchr/testify v1.5.1/go.mod h1:5W2xD1RspED5o8YsWQXVCued0rvSQ+mT+I5cxcmMvtA=
github.com/stretchr/testify v1.6.1/go.mod h1:6Fq8oRcR53rry900zMqJjRRixrwX3KX962/h/Wwjteg=
github.com/stretchr/testify v1.7.0/go.mod h1:6Fq8oRcR53rry900zMqJjRRixrwX3KX962/h/Wwjteg=
github.com/stretchr/testify v1.7.1/go.mod h1:6Fq8oRcR53rry900zMqJjRRixrwX3KX962/h/Wwjteg=
@@ -307,8 +340,9 @@ github.com/tidwall/gjson v1.18.0 h1:FIDeeyB800efLX89e5a8Y0BNH+LOngJyGrIWxG2FKQY=
github.com/tidwall/gjson v1.18.0/go.mod h1:/wbyibRr2FHMks5tjHJ5F8dMZh3AcwJEMf5vlfC0lxk=
github.com/tidwall/match v1.1.1 h1:+Ho715JplO36QYgwN9PGYNhgZvoUSc9X2c80KVTi+GA=
github.com/tidwall/match v1.1.1/go.mod h1:eRSPERbgtNPcGhD8UCthc6PmLEQXEWd3PRB5JTxsfmM=
github.com/tidwall/pretty v1.2.0 h1:RWIZEg2iJ8/g6fDDYzMpobmaoGh5OLl4AXtGUGPcqCs=
github.com/tidwall/pretty v1.2.0/go.mod h1:ITEVvHYasfjBbM0u2Pg8T2nJnzm8xPwvNhhsoaGGjNU=
github.com/tidwall/pretty v1.2.1 h1:qjsOFOWWQl+N3RsoF5/ssm1pHmJJwhjlSbZ51I6wMl4=
github.com/tidwall/pretty v1.2.1/go.mod h1:ITEVvHYasfjBbM0u2Pg8T2nJnzm8xPwvNhhsoaGGjNU=
github.com/tidwall/sjson v1.2.5 h1:kLy8mja+1c9jlljvWTlSazM7cKDRfJuR/bOJhcY5NcY=
github.com/tidwall/sjson v1.2.5/go.mod h1:Fvgq9kS/6ociJEDnK0Fk1cpYF4FIW6ZF7LAe+6jwd28=
github.com/tiktoken-go/tokenizer v0.6.2 h1:t0GN2DvcUZSFWT/62YOgoqb10y7gSXBGs0A+4VCQK+g=
@@ -325,6 +359,10 @@ github.com/ugorji/go/codec v1.1.7/go.mod h1:Ax+UKWsSmolVDwsd+7N3ZtXu+yMGCf907BLY
github.com/ugorji/go/codec v1.2.7/go.mod h1:WGN1fab3R1fzQlVQTkfxVtIBhWDRqOviHU95kRgeqEY=
github.com/ugorji/go/codec v1.2.12 h1:9LC83zGrHhuUA9l16C9AHXAqEV/2wBQ4nkvumAE65EE=
github.com/ugorji/go/codec v1.2.12/go.mod h1:UNopzCgEMSXjBc6AOMqYvWC1ktqTAfzJZUZgYf6w6lg=
github.com/volcengine/volc-sdk-golang v1.0.23 h1:anOslb2Qp6ywnsbyq9jqR0ljuO63kg9PY+4OehIk5R8=
github.com/volcengine/volc-sdk-golang v1.0.23/go.mod h1:AfG/PZRUkHJ9inETvbjNifTDgut25Wbkm2QoYBTbvyU=
github.com/volcengine/volcengine-go-sdk v1.2.10 h1:C227zcVFwQaM0yOpCANP1vGWZbW5ndpnugah6e8czC8=
github.com/volcengine/volcengine-go-sdk v1.2.10/go.mod h1:oxoVo+A17kvkwPkIeIHPVLjSw7EQAm+l/Vau1YGHN+A=
github.com/x448/float16 v0.8.4 h1:qLwI1I70+NjRFUR3zs1JPUCgaCXSh3SW62uAKT1mSBM=
github.com/x448/float16 v0.8.4/go.mod h1:14CWIYCyZA/cWjXOioeEpHeN/83MdbZDRQHoFcYsOfg=
github.com/xyproto/randomstring v1.0.5 h1:YtlWPoRdgMu3NZtP45drfy1GKoojuR7hmRcnhZqKjWU=
@@ -334,6 +372,10 @@ github.com/yapingcat/gomedia v0.0.0-20240906162731-17feea57090c/go.mod h1:WSZ59b
github.com/yuin/goldmark v1.4.13/go.mod h1:6yULJ656Px+3vBD8DxQVa3kxgyrAnzto9xy5taEt/CY=
github.com/yusufpapurcu/wmi v1.2.3 h1:E1ctvB7uKFMOJw3fdOW32DwGE9I7t++CRUEMKvFoFiw=
github.com/yusufpapurcu/wmi v1.2.3/go.mod h1:SBZ9tNy3G9/m5Oi98Zks0QjeHVDvuK0qfxQmPyzfmi0=
gitlab.ecloud.com/ecloud/ecloudsdkcore v1.0.6 h1:b7SFKIjaJBmkXEnih7ArHz3jCff/k56/vOHgP4KeZAI=
gitlab.ecloud.com/ecloud/ecloudsdkcore v1.0.6/go.mod h1:Reukq+A2PUdVHnAQv5oXCJT8nm0Ig1ZO60VfR6bNawM=
gitlab.ecloud.com/ecloud/ecloudsdkmaas v1.0.2 h1:JPmskyGB98qDeWKLrca7ngT7FO6yTiL220KRES5k2/4=
gitlab.ecloud.com/ecloud/ecloudsdkmaas v1.0.2/go.mod h1:UbEEVYgaAUPypTX4Ug721k0scLXZXG8HaaFnDgh2NLI=
go.uber.org/mock v0.6.0 h1:hyF9dfmbgIX5EfOdasqLsWD6xqpNZlXblLB/Dbnwv3Y=
go.uber.org/mock v0.6.0/go.mod h1:KiVJ4BqZJaMj4svdfmHM0AUx4NJYO8ZNpPnZn1Z+BBU=
go.yaml.in/yaml/v3 v3.0.4 h1:tfq32ie2Jv2UxXFdLJdh3jXuOzWiL1fo0bu/FbuKpbc=
@@ -348,10 +390,14 @@ golang.org/x/crypto v0.19.0/go.mod h1:Iy9bg/ha4yyC70EfRS8jz+B6ybOBKMaSxLj6P6oBDf
golang.org/x/crypto v0.23.0/go.mod h1:CKFgDieR+mRhux2Lsu27y0fO304Db0wZe70UKqHu0v8=
golang.org/x/crypto v0.48.0 h1:/VRzVqiRSggnhY7gNRxPauEQ5Drw9haKdM0jqfcCFts=
golang.org/x/crypto v0.48.0/go.mod h1:r0kV5h3qnFPlQnBSrULhlsRfryS2pmewsg+XfMgkVos=
golang.org/x/exp v0.0.0-20190121172915-509febef88a4/go.mod h1:CJ0aWSM057203Lf6IL+f9T1iT9GByDxfZKAQTCR3kQA=
golang.org/x/exp v0.0.0-20250620022241-b7579e27df2b h1:M2rDM6z3Fhozi9O7NWsxAkg/yqS/lQJ6PmkyIV3YP+o=
golang.org/x/exp v0.0.0-20250620022241-b7579e27df2b/go.mod h1:3//PLf8L/X+8b4vuAfHzxeRUl04Adcb341+IGKfnqS8=
golang.org/x/image v0.23.0 h1:HseQ7c2OpPKTPVzNjG5fwJsOTCiiwS4QdsYi5XU6H68=
golang.org/x/image v0.23.0/go.mod h1:wJJBTdLfCCf3tiHa1fNxpZmUI4mmoZvwMCPP0ddoNKY=
golang.org/x/lint v0.0.0-20181026193005-c67002cb31c3/go.mod h1:UVdnD1Gm6xHRNCYTkRU2/jEulfH38KcIWyp/GAMgvoE=
golang.org/x/lint v0.0.0-20190227174305-5b3e6a55c961/go.mod h1:wehouNa3lNwaWXcvxsM5YxQ5yQlVC4a0KAMCusXpPoU=
golang.org/x/lint v0.0.0-20190313153728-d0100b6bd8b3/go.mod h1:6SW0HCj/g11FgYtHlgUYUwCkIfeOF89ocIRzGO/8vkc=
golang.org/x/mod v0.6.0-dev.0.20220419223038-86c51ed26bb4/go.mod h1:jJ57K6gSWd91VN4djpZkiMVwK6gcyfeH4XE8wZrZaV4=
golang.org/x/mod v0.8.0/go.mod h1:iBbtSCu2XBx23ZKBPSOrRkjjQPZFPuis4dIYUhu/chs=
golang.org/x/mod v0.12.0/go.mod h1:iBbtSCu2XBx23ZKBPSOrRkjjQPZFPuis4dIYUhu/chs=
@@ -359,6 +405,10 @@ golang.org/x/mod v0.15.0/go.mod h1:hTbmBsO62+eylJbnUtE2MGJUyE7QWk4xUqPFrRgJ+7c=
golang.org/x/mod v0.17.0/go.mod h1:hTbmBsO62+eylJbnUtE2MGJUyE7QWk4xUqPFrRgJ+7c=
golang.org/x/mod v0.32.0 h1:9F4d3PHLljb6x//jOyokMv3eX+YDeepZSEo3mFJy93c=
golang.org/x/mod v0.32.0/go.mod h1:SgipZ/3h2Ci89DlEtEXWUk/HteuRin+HHhN+WbNhguU=
golang.org/x/net v0.0.0-20180724234803-3673e40ba225/go.mod h1:mL1N/T3taQHkDXs73rZJwtUhF3w3ftmwwsq0BUmARs4=
golang.org/x/net v0.0.0-20180826012351-8a410e7b638d/go.mod h1:mL1N/T3taQHkDXs73rZJwtUhF3w3ftmwwsq0BUmARs4=
golang.org/x/net v0.0.0-20190213061140-3a22650c66bd/go.mod h1:mL1N/T3taQHkDXs73rZJwtUhF3w3ftmwwsq0BUmARs4=
golang.org/x/net v0.0.0-20190311183353-d8887717615a/go.mod h1:t9HGtf8HONx5eT2rtn7q6eTqICYqUVnKs3thJo3Qplg=
golang.org/x/net v0.0.0-20190620200207-3b0461eec859/go.mod h1:z5CRVTTTmAJ677TzLLGU+0bjPO0LkuOLi4/5GtJWs/s=
golang.org/x/net v0.0.0-20210226172049-e18ecbb05110/go.mod h1:m0MpNAwzfU5UDzcl9v0D8zg8gWTRqZa9RBIspLL5mdg=
golang.org/x/net v0.0.0-20210520170846-37e1c6afe023/go.mod h1:9nx3DQGgdP8bBQD5qxJ1jj9UTztislL4KSBs9R2vV5Y=
@@ -370,6 +420,9 @@ golang.org/x/net v0.21.0/go.mod h1:bIjVDfnllIU7BJ2DNgfnXvpSvtn8VRwhlsaeUTyUS44=
golang.org/x/net v0.25.0/go.mod h1:JkAGAh7GEvH74S6FOH42FLoXpXbE/aqXSrIQjXgsiwM=
golang.org/x/net v0.49.0 h1:eeHFmOGUTtaaPSGNmjBKpbng9MulQsJURQUAfUwY++o=
golang.org/x/net v0.49.0/go.mod h1:/ysNB2EvaqvesRkuLAyjI1ycPZlQHM3q01F02UY/MV8=
golang.org/x/oauth2 v0.0.0-20180821212333-d2e6202438be/go.mod h1:N/0e6XlmueqKjAGxoOufVs8QHGRruUQn6yWY3a++T0U=
golang.org/x/sync v0.0.0-20180314180146-1d60e4601c6f/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
golang.org/x/sync v0.0.0-20181108010431-42b317875d0f/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
golang.org/x/sync v0.0.0-20190423024810-112230192c58/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
golang.org/x/sync v0.0.0-20220722155255-886fb9371eb4/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
golang.org/x/sync v0.1.0/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
@@ -379,6 +432,7 @@ golang.org/x/sync v0.7.0/go.mod h1:Czt+wKu1gCyEFDUtn0jG5QVvpJ6rzVqr5aXyt9drQfk=
golang.org/x/sync v0.10.0/go.mod h1:Czt+wKu1gCyEFDUtn0jG5QVvpJ6rzVqr5aXyt9drQfk=
golang.org/x/sync v0.19.0 h1:vV+1eWNmZ5geRlYjzm2adRgW2/mcpevXNg50YZtPCE4=
golang.org/x/sync v0.19.0/go.mod h1:9KTHXmSnoGruLpwFjVSX0lNNA75CykiMECbovNTZqGI=
golang.org/x/sys v0.0.0-20180830151530-49385e6e1522/go.mod h1:STP8DvDyc/dI5b8T5hshtkjS+E42TnysNCUPdjciGhY=
golang.org/x/sys v0.0.0-20190215142949-d0b11bdaac8a/go.mod h1:STP8DvDyc/dI5b8T5hshtkjS+E42TnysNCUPdjciGhY=
golang.org/x/sys v0.0.0-20190726091711-fc99dfbffb4e/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs=
golang.org/x/sys v0.0.0-20190916202348-b4ddaad3f8a3/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs=
@@ -421,6 +475,10 @@ golang.org/x/text v0.21.0/go.mod h1:4IBbMaMmOPCJ8SecivzSH54+73PCFmPWxNTLm+vZkEQ=
golang.org/x/text v0.34.0 h1:oL/Qq0Kdaqxa1KbNeMKwQq0reLCCaFtqu2eNuSeNHbk=
golang.org/x/text v0.34.0/go.mod h1:homfLqTYRFyVYemLBFl5GgL/DWEiH5wcsQ5gSh1yziA=
golang.org/x/tools v0.0.0-20180917221912-90fa682c2a6e/go.mod h1:n7NCudcB/nEzxVGmLbDWY5pfWTLqBcC2KZ6jyYvM4mQ=
golang.org/x/tools v0.0.0-20190114222345-bf090417da8b/go.mod h1:n7NCudcB/nEzxVGmLbDWY5pfWTLqBcC2KZ6jyYvM4mQ=
golang.org/x/tools v0.0.0-20190226205152-f727befe758c/go.mod h1:9Yl7xja0Znq3iFh3HoIrodX9oNMXvdceNzlUR8zjMvY=
golang.org/x/tools v0.0.0-20190311212946-11955173bddd/go.mod h1:LCzVGOaR6xXOjkQ3onu1FJEFr0SW1gC7cKk1uF8kGRs=
golang.org/x/tools v0.0.0-20190524140312-2c0ae7006135/go.mod h1:RgjU9mgBXZiqYHBnxXauZ1Gv1EHHAz9KjViQ78xBX0Q=
golang.org/x/tools v0.0.0-20191119224855-298f0cb1881e/go.mod h1:b+2E5dAYhXwXZwtnZ6UAqBI28+e2cm9otk0dWdXHAEo=
golang.org/x/tools v0.1.12/go.mod h1:hNGJHUnrk76NpqgfD5Aqm5Crs+Hm0VOH/i9J2+nxYbc=
golang.org/x/tools v0.6.0/go.mod h1:Xwgl3UAJ/d3gWutnCtw505GrjyAbvKui8lOU390QaIU=
@@ -430,12 +488,31 @@ golang.org/x/tools v0.41.0 h1:a9b8iMweWG+S0OBnlU36rzLp20z1Rp10w+IY2czHTQc=
golang.org/x/tools v0.41.0/go.mod h1:XSY6eDqxVNiYgezAVqqCeihT4j1U2CCsqvH3WhQpnlg=
golang.org/x/xerrors v0.0.0-20190717185122-a985d3407aa7/go.mod h1:I/5z698sn9Ka8TeJc9MKroUUfqBBauWjQqLJ2OPfmY0=
golang.org/x/xerrors v0.0.0-20191204190536-9bdfabe68543/go.mod h1:I/5z698sn9Ka8TeJc9MKroUUfqBBauWjQqLJ2OPfmY0=
google.golang.org/appengine v1.1.0/go.mod h1:EbEs0AVv82hx2wNQdGPgUI5lhzA/G0D9YwlJXL52JkM=
google.golang.org/appengine v1.4.0/go.mod h1:xpcJRLb0r/rnEns0DIKYYv+WjYCduHsrkT7/EB5XEv4=
google.golang.org/genproto v0.0.0-20180817151627-c66870c02cf8/go.mod h1:JiN7NxoALGmiZfu7CAH4rXhgtRTLTxftemlI0sWmxmc=
google.golang.org/genproto v0.0.0-20190819201941-24fa4b261c55/go.mod h1:DMBHOl98Agz4BDEuKkezgsaosCRResVns1a3J2ZsMNc=
google.golang.org/genproto v0.0.0-20200526211855-cb27e3aa2013/go.mod h1:NbSheEEYHJ7i3ixzK3sjbqSGDJWnxyFXZblF3eUsNvo=
google.golang.org/grpc v1.19.0/go.mod h1:mqu4LbDTu4XGKhr4mRzUsmM4RtVoemTSY81AxZiDr8c=
google.golang.org/grpc v1.23.0/go.mod h1:Y5yQAOtifL1yxbo5wqy6BxZv8vAUGQwXBOALyacEbxg=
google.golang.org/grpc v1.27.0/go.mod h1:qbnxyOmOxrQa7FizSgH+ReBfzJrCY1pSN7KXBS8abTk=
google.golang.org/protobuf v0.0.0-20200109180630-ec00e32a8dfd/go.mod h1:DFci5gLYBciE7Vtevhsrf46CRTquxDuWsQurQQe4oz8=
google.golang.org/protobuf v0.0.0-20200221191635-4d8936d0db64/go.mod h1:kwYJMbMJ01Woi6D6+Kah6886xMZcty6N08ah7+eCXa0=
google.golang.org/protobuf v0.0.0-20200228230310-ab0ca4ff8a60/go.mod h1:cfTl7dwQJ+fmap5saPgwCLgHXTUD7jkjRqWcaiX5VyM=
google.golang.org/protobuf v1.20.1-0.20200309200217-e05f789c0967/go.mod h1:A+miEFZTKqfCUM6K7xSMQL9OKL/b6hQv+e19PK+JZNE=
google.golang.org/protobuf v1.21.0/go.mod h1:47Nbq4nVaFHyn7ilMalzfO3qCViNmqZ2kzikPIcrTAo=
google.golang.org/protobuf v1.22.0/go.mod h1:EGpADcykh3NcUnDUJcl1+ZksZNG86OlYog2l/sGQquU=
google.golang.org/protobuf v1.23.0/go.mod h1:EGpADcykh3NcUnDUJcl1+ZksZNG86OlYog2l/sGQquU=
google.golang.org/protobuf v1.23.1-0.20200526195155-81db48ad09cc/go.mod h1:EGpADcykh3NcUnDUJcl1+ZksZNG86OlYog2l/sGQquU=
google.golang.org/protobuf v1.25.0/go.mod h1:9JNX74DMeImyA3h4bdi1ymwjUzf21/xIlbajtzgsN7c=
google.golang.org/protobuf v1.26.0-rc.1/go.mod h1:jlhhOSvTdKEhbULTjvd4ARK9grFBp09yW+WbY/TyQbw=
google.golang.org/protobuf v1.28.0/go.mod h1:HV8QOd/L58Z+nl8r43ehVNZIU/HEI6OcFqwMG9pJV4I=
google.golang.org/protobuf v1.31.0/go.mod h1:HV8QOd/L58Z+nl8r43ehVNZIU/HEI6OcFqwMG9pJV4I=
google.golang.org/protobuf v1.36.5 h1:tPhr+woSbjfYvY6/GPufUoYizxw1cF/yFoxJ2fmpwlM=
google.golang.org/protobuf v1.36.5/go.mod h1:9fA7Ob0pmnwhb644+1+CVWFRbNajQ6iRojtC/QF5bRE=
gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405/go.mod h1:Co6ibVJAznAaIkqp8huTwlJQCZ016jof/cbN4VW5Yz0=
gopkg.in/check.v1 v1.0.0-20180628173108-788fd7840127/go.mod h1:Co6ibVJAznAaIkqp8huTwlJQCZ016jof/cbN4VW5Yz0=
gopkg.in/check.v1 v1.0.0-20190902080502-41f04d3bba15/go.mod h1:Co6ibVJAznAaIkqp8huTwlJQCZ016jof/cbN4VW5Yz0=
gopkg.in/check.v1 v1.0.0-20201130134442-10cb98267c6c h1:Hei/4ADfdWqJk1ZMxUNpqntNwaWcugrBjAiHlqqRiVk=
gopkg.in/check.v1 v1.0.0-20201130134442-10cb98267c6c/go.mod h1:JHkPIbrfpd72SG/EVd6muEfDQjcINNoR0C8j2r3qZ4Q=
gopkg.in/errgo.v2 v2.1.0/go.mod h1:hNsd1EY+bozCKY1Ytp96fpM3vjJbqLJn88ws8XvfDNI=
@@ -458,6 +535,8 @@ gorm.io/driver/postgres v1.5.2/go.mod h1:fmpX0m2I1PKuR7mKZiEluwrP3hbs+ps7JIGMUBp
gorm.io/gorm v1.23.8/go.mod h1:l2lP/RyAtc1ynaTjFksBde/O8v9oOGIApu2/xRitmZk=
gorm.io/gorm v1.25.2 h1:gs1o6Vsa+oVKG/a9ElL3XgyGfghFfkKA2SInQaCyMho=
gorm.io/gorm v1.25.2/go.mod h1:L4uxeKpfBml98NYqVqwAdmV1a2nBtAec/cf3fpucW/k=
honnef.co/go/tools v0.0.0-20190102054323-c2f93a96b099/go.mod h1:rf3lG4BRIbNafJWhAfAdb/ePZxsR/4RtNHQocxwk9r4=
honnef.co/go/tools v0.0.0-20190523083050-ea95bdfd59fc/go.mod h1:rf3lG4BRIbNafJWhAfAdb/ePZxsR/4RtNHQocxwk9r4=
modernc.org/cc/v4 v4.26.5 h1:xM3bX7Mve6G8K8b+T11ReenJOT+BmVqQj0FY5T4+5Y4=
modernc.org/cc/v4 v4.26.5/go.mod h1:uVtb5OGqUKpoLWhqwNQo/8LwvoiEBLvZXIQ/SmO6mL0=
modernc.org/ccgo/v4 v4.28.1 h1:wPKYn5EC/mYTqBO373jKjvX2n+3+aK7+sICCv4Fjy1A=


+ 64
- 0
middleware/gzip_test.go Просмотреть файл

@@ -0,0 +1,64 @@
package middleware

import (
"bytes"
"compress/gzip"
"io"
"net/http"
"net/http/httptest"
"testing"

"github.com/andybalholm/brotli"
"github.com/gin-gonic/gin"
"github.com/stretchr/testify/require"
)

func TestDecompressRequestMiddlewareExposesGzipAndBrotliPayload(t *testing.T) {
gin.SetMode(gin.TestMode)
for _, encoding := range []string{"gzip", "br"} {
t.Run(encoding, func(t *testing.T) {
var compressed bytes.Buffer
var writer io.WriteCloser
if encoding == "gzip" {
writer = gzip.NewWriter(&compressed)
} else {
writer = brotli.NewWriter(&compressed)
}
_, err := writer.Write([]byte(`{"message":"hello"}`))
require.NoError(t, err)
require.NoError(t, writer.Close())

router := gin.New()
router.Use(DecompressRequestMiddleware())
router.POST("/", func(c *gin.Context) {
body, err := io.ReadAll(c.Request.Body)
require.NoError(t, err)
require.Equal(t, `{"message":"hello"}`, string(body))
require.Empty(t, c.GetHeader("Content-Encoding"))
c.Status(http.StatusNoContent)
})

request := httptest.NewRequest(http.MethodPost, "/", &compressed)
request.Header.Set("Content-Encoding", encoding)
response := httptest.NewRecorder()
router.ServeHTTP(response, request)
require.Equal(t, http.StatusNoContent, response.Code)
})
}
}

func TestDecompressRequestMiddlewareRejectsMalformedGzip(t *testing.T) {
gin.SetMode(gin.TestMode)
reachedHandler := false
router := gin.New()
router.Use(DecompressRequestMiddleware())
router.POST("/", func(c *gin.Context) { reachedHandler = true })

request := httptest.NewRequest(http.MethodPost, "/", bytes.NewBufferString("not-gzip"))
request.Header.Set("Content-Encoding", "gzip")
response := httptest.NewRecorder()
router.ServeHTTP(response, request)

require.Equal(t, http.StatusBadRequest, response.Code)
require.False(t, reachedHandler)
}

+ 44
- 0
middleware/http_headers_test.go Просмотреть файл

@@ -0,0 +1,44 @@
package middleware

import (
"net/http"
"net/http/httptest"
"testing"

"github.com/QuantumNous/new-api/common"
"github.com/gin-gonic/gin"
"github.com/stretchr/testify/require"
)

func TestDisableCacheSetsNoCacheHeaders(t *testing.T) {
gin.SetMode(gin.TestMode)
router := gin.New()
router.Use(DisableCache())
router.GET("/", func(c *gin.Context) { c.Status(http.StatusNoContent) })

response := httptest.NewRecorder()
router.ServeHTTP(response, httptest.NewRequest(http.MethodGet, "/", nil))

require.Equal(t, http.StatusNoContent, response.Code)
require.Equal(t, "no-store, no-cache, must-revalidate, private, max-age=0", response.Header().Get("Cache-Control"))
require.Equal(t, "no-cache", response.Header().Get("Pragma"))
require.Equal(t, "0", response.Header().Get("Expires"))
}

func TestRequestIdPropagatesSameIDToContextAndResponse(t *testing.T) {
gin.SetMode(gin.TestMode)
router := gin.New()
router.Use(RequestId())
router.GET("/", func(c *gin.Context) {
id := c.GetString(common.RequestIdKey)
require.NotEmpty(t, id)
require.Equal(t, id, c.Request.Context().Value(common.RequestIdKey))
c.Status(http.StatusNoContent)
})

response := httptest.NewRecorder()
router.ServeHTTP(response, httptest.NewRequest(http.MethodGet, "/", nil))

require.Equal(t, http.StatusNoContent, response.Code)
require.NotEmpty(t, response.Header().Get(common.RequestIdKey))
}

+ 41
- 0
middleware/recover_test.go Просмотреть файл

@@ -0,0 +1,41 @@
package middleware

import (
"net/http"
"net/http/httptest"
"testing"

"github.com/QuantumNous/new-api/common"
"github.com/gin-gonic/gin"
"github.com/stretchr/testify/require"
)

func TestRelayPanicRecoverConvertsPanicToServerError(t *testing.T) {
gin.SetMode(gin.TestMode)
router := gin.New()
router.Use(RelayPanicRecover())
router.GET("/", func(c *gin.Context) { panic("upstream exploded") })

response := httptest.NewRecorder()
router.ServeHTTP(response, httptest.NewRequest(http.MethodGet, "/", nil))

require.Equal(t, http.StatusInternalServerError, response.Code)
require.Contains(t, response.Body.String(), "new_api_panic")
require.Contains(t, response.Body.String(), "upstream exploded")
}

func TestTurnstileCheckAllowsRequestWhenFeatureIsDisabled(t *testing.T) {
oldEnabled := common.TurnstileCheckEnabled
common.TurnstileCheckEnabled = false
t.Cleanup(func() { common.TurnstileCheckEnabled = oldEnabled })

gin.SetMode(gin.TestMode)
router := gin.New()
router.Use(TurnstileCheck())
router.GET("/", func(c *gin.Context) { c.Status(http.StatusNoContent) })

response := httptest.NewRecorder()
router.ServeHTTP(response, httptest.NewRequest(http.MethodGet, "/", nil))

require.Equal(t, http.StatusNoContent, response.Code)
}

+ 86
- 16
model/channel.go Просмотреть файл

@@ -56,6 +56,10 @@ type Channel struct {

// cache info
Keys []string `json:"-" gorm:"-"`

// Asset credential metadata is populated only for management API responses.
AssetCredentialConfigured bool `json:"asset_credential_configured,omitempty" gorm:"-"`
AssetCredentialPoolID string `json:"asset_credential_pool_id,omitempty" gorm:"-"`
}

type ChannelInfo struct {
@@ -283,7 +287,6 @@ func GetAllChannels(startIdx int, num int, selectAll bool, idSort bool) ([]*Chan
return channels, err
}


func GetChannelsByTag(tag string, idSort bool, selectAll bool) ([]*Channel, error) {
var channels []*Channel
order := "priority desc"
@@ -403,7 +406,7 @@ func BatchDeleteChannels(ids []int) error {
return tx.Error
}
for _, chunk := range lo.Chunk(ids, 200) {
if err := tx.Where("id in (?)", chunk).Delete(&Channel{}).Error; err != nil {
if err := DeleteChannelAssetCredentialsWithTx(tx, chunk); err != nil {
tx.Rollback()
return err
}
@@ -411,6 +414,10 @@ func BatchDeleteChannels(ids []int) error {
tx.Rollback()
return err
}
if err := tx.Where("id in (?)", chunk).Delete(&Channel{}).Error; err != nil {
tx.Rollback()
return err
}
}
return tx.Commit().Error
}
@@ -465,6 +472,12 @@ func (channel *Channel) Insert() error {
}

func (channel *Channel) Update() error {
return DB.Transaction(func(tx *gorm.DB) error {
return channel.UpdateWithTx(tx)
})
}

func (channel *Channel) UpdateWithTx(tx *gorm.DB) error {
// If this is a multi-key channel, recalculate MultiKeySize based on the current key list to avoid inconsistency after editing keys
if channel.ChannelInfo.IsMultiKey {
var keyStr string
@@ -472,7 +485,8 @@ func (channel *Channel) Update() error {
keyStr = channel.Key
} else {
// If key is not provided, read the existing key from the database
if existing, err := GetChannelById(channel.Id, true); err == nil {
var existing Channel
if err := tx.First(&existing, "id = ?", channel.Id).Error; err == nil {
keyStr = existing.Key
}
}
@@ -503,16 +517,30 @@ func (channel *Channel) Update() error {
}
}
}
var err error
err = DB.Model(channel).Updates(channel).Error
err := tx.Model(channel).Updates(channel).Error
if err != nil {
return err
}
DB.Model(channel).First(channel, "id = ?", channel.Id)
err = channel.UpdateAbilities(nil)
if err = tx.First(channel, "id = ?", channel.Id).Error; err != nil {
return err
}
err = channel.UpdateAbilities(tx)
return err
}

func InsertChannelWithAssetCredential(channel *Channel, credential *ChannelAssetCredential) error {
return DB.Transaction(func(tx *gorm.DB) error {
if err := tx.Create(channel).Error; err != nil {
return err
}
if err := channel.AddAbilities(tx); err != nil {
return err
}
credential.ChannelId = channel.Id
return UpsertChannelAssetCredentialWithTx(tx, credential)
})
}

func (channel *Channel) UpdateResponseTime(responseTime int64) {
err := DB.Model(channel).Select("response_time", "test_time").Updates(Channel{
TestTime: common.GetTimestamp(),
@@ -534,13 +562,19 @@ func (channel *Channel) UpdateBalance(balance float64) {
}

func (channel *Channel) Delete() error {
var err error
err = DB.Delete(channel).Error
if err != nil {
return DB.Transaction(func(tx *gorm.DB) error {
return channel.DeleteWithTx(tx)
})
}

func (channel *Channel) DeleteWithTx(tx *gorm.DB) error {
if err := DeleteChannelAssetCredentialWithTx(tx, channel.Id); err != nil {
return err
}
err = channel.DeleteAbilities()
return err
if err := tx.Where("channel_id = ?", channel.Id).Delete(&Ability{}).Error; err != nil {
return err
}
return tx.Delete(channel).Error
}

var channelStatusLock sync.Mutex
@@ -778,13 +812,49 @@ func updateChannelUsedQuota(id int, quota int) {
}

func DeleteChannelByStatus(status int64) (int64, error) {
result := DB.Where("status = ?", status).Delete(&Channel{})
return result.RowsAffected, result.Error
var ids []int
if err := DB.Model(&Channel{}).Where("status = ?", status).Pluck("id", &ids).Error; err != nil {
return 0, err
}
if len(ids) == 0 {
return 0, nil
}
var rows int64
err := DB.Transaction(func(tx *gorm.DB) error {
if err := DeleteChannelAssetCredentialsWithTx(tx, ids); err != nil {
return err
}
if err := tx.Where("channel_id IN ?", ids).Delete(&Ability{}).Error; err != nil {
return err
}
result := tx.Where("id IN ?", ids).Delete(&Channel{})
rows = result.RowsAffected
return result.Error
})
return rows, err
}

func DeleteDisabledChannel() (int64, error) {
result := DB.Where("status = ? or status = ?", common.ChannelStatusAutoDisabled, common.ChannelStatusManuallyDisabled).Delete(&Channel{})
return result.RowsAffected, result.Error
var ids []int
if err := DB.Model(&Channel{}).Where("status = ? or status = ?", common.ChannelStatusAutoDisabled, common.ChannelStatusManuallyDisabled).Pluck("id", &ids).Error; err != nil {
return 0, err
}
if len(ids) == 0 {
return 0, nil
}
var rows int64
err := DB.Transaction(func(tx *gorm.DB) error {
if err := DeleteChannelAssetCredentialsWithTx(tx, ids); err != nil {
return err
}
if err := tx.Where("channel_id IN ?", ids).Delete(&Ability{}).Error; err != nil {
return err
}
result := tx.Where("id IN ?", ids).Delete(&Channel{})
rows = result.RowsAffected
return result.Error
})
return rows, err
}

func GetPaginatedTags(offset int, limit int) ([]*string, error) {


+ 99
- 0
model/channel_asset_credential.go Просмотреть файл

@@ -0,0 +1,99 @@
package model

import (
"github.com/QuantumNous/new-api/common"
"gorm.io/gorm"
"gorm.io/gorm/clause"
)

// ChannelAssetCredential stores the China Mobile asset credentials separately
// from the channel video-generation key.
type ChannelAssetCredential struct {
Id int `json:"id"`
ChannelId int `json:"channel_id" gorm:"uniqueIndex;not null"`
AccessKey string `json:"-" gorm:"not null;size:255"`
SecretKey string `json:"-" gorm:"not null;size:255"`
PoolID string `json:"pool_id" gorm:"size:255"`
CreatedAt int64 `json:"created_at" gorm:"bigint;not null"`
UpdatedAt int64 `json:"updated_at" gorm:"bigint;not null"`
}

// ChannelAssetCredentialSummary is safe to expose in channel management APIs.
type ChannelAssetCredentialSummary struct {
ChannelId int
PoolID string
}

func GetChannelAssetCredential(channelID int) (*ChannelAssetCredential, error) {
var credential ChannelAssetCredential
err := DB.Where("channel_id = ?", channelID).First(&credential).Error
if err == gorm.ErrRecordNotFound {
return nil, nil
}
if err != nil {
return nil, err
}
return &credential, nil
}

func UpsertChannelAssetCredential(credential *ChannelAssetCredential) error {
return DB.Transaction(func(tx *gorm.DB) error {
return UpsertChannelAssetCredentialWithTx(tx, credential)
})
}

func UpsertChannelAssetCredentialWithTx(tx *gorm.DB, credential *ChannelAssetCredential) error {
now := common.GetTimestamp()
credential.CreatedAt = now
credential.UpdatedAt = now
return tx.Clauses(clause.OnConflict{
Columns: []clause.Column{{Name: "channel_id"}},
DoUpdates: clause.Assignments(map[string]any{
"access_key": credential.AccessKey,
"secret_key": credential.SecretKey,
"pool_id": credential.PoolID,
"updated_at": credential.UpdatedAt,
}),
}).Create(credential).Error
}

func DeleteChannelAssetCredentialWithTx(tx *gorm.DB, channelID int) error {
return tx.Where("channel_id = ?", channelID).Delete(&ChannelAssetCredential{}).Error
}

func DeleteChannelAssetCredential(channelID int) error {
return DB.Transaction(func(tx *gorm.DB) error {
return DeleteChannelAssetCredentialWithTx(tx, channelID)
})
}

func DeleteChannelAssetCredentialsWithTx(tx *gorm.DB, channelIDs []int) error {
if len(channelIDs) == 0 {
return nil
}
return tx.Where("channel_id IN ?", channelIDs).Delete(&ChannelAssetCredential{}).Error
}

func DeleteChannelAssetCredentials(channelIDs []int) error {
return DB.Transaction(func(tx *gorm.DB) error {
return DeleteChannelAssetCredentialsWithTx(tx, channelIDs)
})
}

func GetChannelAssetCredentialSummaries(channelIDs []int) (map[int]ChannelAssetCredentialSummary, error) {
summaries := make(map[int]ChannelAssetCredentialSummary)
if len(channelIDs) == 0 {
return summaries, nil
}
var rows []ChannelAssetCredentialSummary
if err := DB.Model(&ChannelAssetCredential{}).
Select("channel_id", "pool_id").
Where("channel_id IN ?", channelIDs).
Find(&rows).Error; err != nil {
return nil, err
}
for _, row := range rows {
summaries[row.ChannelId] = row
}
return summaries, nil
}

+ 132
- 0
model/channel_asset_credential_test.go Просмотреть файл

@@ -0,0 +1,132 @@
package model

import (
"testing"

"github.com/glebarez/sqlite"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"gorm.io/gorm"
"gorm.io/gorm/logger"
)

func setupChannelAssetCredentialDB(t *testing.T) *gorm.DB {
t.Helper()

db, err := gorm.Open(sqlite.Open("file:channel_asset_credentials?mode=memory&cache=shared"), &gorm.Config{
Logger: logger.Default.LogMode(logger.Silent),
})
require.NoError(t, err)
sqlDB, err := db.DB()
require.NoError(t, err)
sqlDB.SetMaxOpenConns(1)

originalDB := DB
DB = db
t.Cleanup(func() {
DB = originalDB
require.NoError(t, sqlDB.Close())
})
return db
}

func TestChannelAssetCredentialUpsertKeepsOneCredentialPerChannel(t *testing.T) {
db := setupChannelAssetCredentialDB(t)
require.NoError(t, db.AutoMigrate(&ChannelAssetCredential{}))

require.NoError(t, UpsertChannelAssetCredential(&ChannelAssetCredential{
ChannelId: 61,
AccessKey: "ak-old",
SecretKey: "sk-old",
PoolID: "pool-old",
}))
require.NoError(t, UpsertChannelAssetCredential(&ChannelAssetCredential{
ChannelId: 61,
AccessKey: "ak-new",
SecretKey: "sk-new",
PoolID: "pool-new",
}))

credential, err := GetChannelAssetCredential(61)
require.NoError(t, err)
require.NotNil(t, credential)
assert.Equal(t, "ak-new", credential.AccessKey)
assert.Equal(t, "sk-new", credential.SecretKey)
assert.Equal(t, "pool-new", credential.PoolID)

var count int64
require.NoError(t, db.Model(&ChannelAssetCredential{}).Where("channel_id = ?", 61).Count(&count).Error)
assert.Equal(t, int64(1), count)
}

func TestChannelAssetCredentialSummariesDoNotContainSecrets(t *testing.T) {
db := setupChannelAssetCredentialDB(t)
require.NoError(t, db.AutoMigrate(&ChannelAssetCredential{}))
require.NoError(t, UpsertChannelAssetCredential(&ChannelAssetCredential{
ChannelId: 61,
AccessKey: "ak-secret",
SecretKey: "sk-secret",
PoolID: "pool-61",
}))

summaries, err := GetChannelAssetCredentialSummaries([]int{61, 62})
require.NoError(t, err)
require.Contains(t, summaries, 61)
assert.Equal(t, "pool-61", summaries[61].PoolID)
assert.NotContains(t, summaries, 62)
}

func TestChannelDeleteRemovesOnlyItsAssetCredential(t *testing.T) {
db := setupChannelAssetCredentialDB(t)
require.NoError(t, db.AutoMigrate(&Channel{}, &Ability{}, &ChannelAssetCredential{}))
for _, id := range []int{61, 62} {
require.NoError(t, db.Create(&Channel{Id: id, Key: "key", Name: "channel", Group: "default", Models: "model"}).Error)
require.NoError(t, UpsertChannelAssetCredential(&ChannelAssetCredential{ChannelId: id, AccessKey: "ak", SecretKey: "sk"}))
}

require.NoError(t, (&Channel{Id: 61}).Delete())
credential, err := GetChannelAssetCredential(61)
require.NoError(t, err)
assert.Nil(t, credential)
credential, err = GetChannelAssetCredential(62)
require.NoError(t, err)
assert.NotNil(t, credential)
}

func TestDeleteDisabledChannelRemovesAssetCredentials(t *testing.T) {
db := setupChannelAssetCredentialDB(t)
require.NoError(t, db.AutoMigrate(&Channel{}, &Ability{}, &ChannelAssetCredential{}))
require.NoError(t, db.Create(&Channel{Id: 61, Key: "key", Name: "disabled", Group: "default", Models: "model", Status: 2}).Error)
require.NoError(t, db.Create(&Channel{Id: 62, Key: "key", Name: "enabled", Group: "default", Models: "model", Status: 1}).Error)
for _, id := range []int{61, 62} {
require.NoError(t, UpsertChannelAssetCredential(&ChannelAssetCredential{ChannelId: id, AccessKey: "ak", SecretKey: "sk"}))
}

_, err := DeleteDisabledChannel()
require.NoError(t, err)
credential, err := GetChannelAssetCredential(61)
require.NoError(t, err)
assert.Nil(t, credential)
credential, err = GetChannelAssetCredential(62)
require.NoError(t, err)
assert.NotNil(t, credential)
}

func TestDeleteChannelByStatusRemovesAssetCredentials(t *testing.T) {
db := setupChannelAssetCredentialDB(t)
require.NoError(t, db.AutoMigrate(&Channel{}, &Ability{}, &ChannelAssetCredential{}))
require.NoError(t, db.Create(&Channel{Id: 61, Key: "key", Name: "disabled", Group: "default", Models: "model", Status: 2}).Error)
require.NoError(t, db.Create(&Channel{Id: 62, Key: "key", Name: "enabled", Group: "default", Models: "model", Status: 1}).Error)
for _, id := range []int{61, 62} {
require.NoError(t, UpsertChannelAssetCredential(&ChannelAssetCredential{ChannelId: id, AccessKey: "ak", SecretKey: "sk"}))
}

_, err := DeleteChannelByStatus(2)
require.NoError(t, err)
credential, err := GetChannelAssetCredential(61)
require.NoError(t, err)
assert.Nil(t, credential)
credential, err = GetChannelAssetCredential(62)
require.NoError(t, err)
assert.NotNil(t, credential)
}

+ 44
- 0
model/custom_oauth_provider_test.go Просмотреть файл

@@ -0,0 +1,44 @@
package model

import (
"testing"

"github.com/stretchr/testify/require"
)

func validCustomOAuthProvider() *CustomOAuthProvider {
return &CustomOAuthProvider{
Name: "Example OAuth",
Slug: "Example-Provider",
ClientId: "client-id",
AuthorizationEndpoint: "https://id.example/authorize",
TokenEndpoint: "https://id.example/token",
UserInfoEndpoint: "https://id.example/userinfo",
}
}

func TestValidateCustomOAuthProviderNormalizesSlugAndAppliesDefaults(t *testing.T) {
provider := validCustomOAuthProvider()

err := validateCustomOAuthProvider(provider)

require.NoError(t, err)
require.Equal(t, "example-provider", provider.Slug)
require.Equal(t, "sub", provider.UserIdField)
require.Equal(t, "preferred_username", provider.UsernameField)
require.Equal(t, "openid profile email", provider.Scopes)
}

func TestValidateCustomOAuthProviderRejectsInvalidSlugAndPolicy(t *testing.T) {
invalidSlug := validCustomOAuthProvider()
invalidSlug.Slug = "bad_slug"
require.ErrorContains(t, validateCustomOAuthProvider(invalidSlug), "slug")

unsupportedOp := validCustomOAuthProvider()
unsupportedOp.AccessPolicy = `{"conditions":[{"field":"role","op":"matches","value":"admin"}]}`
require.ErrorContains(t, validateCustomOAuthProvider(unsupportedOp), "unsupported")

nonArrayMembership := validCustomOAuthProvider()
nonArrayMembership.AccessPolicy = `{"conditions":[{"field":"role","op":"in","value":"admin"}]}`
require.ErrorContains(t, validateCustomOAuthProvider(nonArrayMembership), "must be an array")
}

+ 73
- 0
model/log.go Просмотреть файл

@@ -4,6 +4,7 @@ import (
"context"
"errors"
"fmt"
"strings"
"time"

"github.com/QuantumNous/new-api/common"
@@ -42,6 +43,69 @@ type Log struct {
Other string `json:"other"`
}

type taskBillingLogOther struct {
BillingMode string `json:"billing_mode"`
BillingPhase string `json:"billing_phase"`
TaskID string `json:"task_id"`
}

// enrichTaskBillingLogs adds terminal billing records for Matrix preconsume logs
// that landed on a different raw log page.
func enrichTaskBillingLogs(logs []*Log) ([]*Log, error) {
taskIDs := make(map[string]struct{})
userIDs := make(map[int]struct{})
existingIDs := make(map[int]struct{}, len(logs))
for _, log := range logs {
existingIDs[log.Id] = struct{}{}
var other taskBillingLogOther
if err := common.UnmarshalJsonStr(log.Other, &other); err != nil {
continue
}
if other.BillingMode == "matrix" && other.BillingPhase == "preconsume" && other.TaskID != "" {
taskIDs[other.TaskID] = struct{}{}
userIDs[log.UserId] = struct{}{}
}
}
if len(taskIDs) == 0 {
return logs, nil
}

patterns := make([]string, 0, len(taskIDs))
for taskID := range taskIDs {
patterns = append(patterns, "%\"task_id\":\""+taskID+"\"%")
}
users := make([]int, 0, len(userIDs))
for userID := range userIDs {
users = append(users, userID)
}
conditions := make([]string, len(patterns))
args := make([]any, len(patterns))
for i, pattern := range patterns {
conditions[i] = "other LIKE ?"
args[i] = pattern
}
tx := LOG_DB.Where("type IN ? AND user_id IN ?", []int{LogTypeConsume, LogTypeRefund}, users).
Where("("+strings.Join(conditions, " OR ")+")", args...)
var candidates []*Log
if err := tx.Find(&candidates).Error; err != nil {
return nil, err
}
for _, log := range candidates {
if _, exists := existingIDs[log.Id]; exists {
continue
}
var other taskBillingLogOther
if err := common.UnmarshalJsonStr(log.Other, &other); err != nil {
continue
}
if _, wanted := taskIDs[other.TaskID]; wanted && other.BillingMode == "matrix" &&
(other.BillingPhase == "settlement" || other.BillingPhase == "refund") {
logs = append(logs, log)
}
}
return logs, nil
}

// don't use iota, avoid change log type value
const (
LogTypeUnknown = 0
@@ -301,6 +365,10 @@ func GetAllLogs(logType int, startTimestamp int64, endTimestamp int64, modelName
if err != nil {
return nil, 0, err
}
logs, err = enrichTaskBillingLogs(logs)
if err != nil {
return nil, 0, err
}

channelIds := types.NewSet[int]()
for _, log := range logs {
@@ -394,6 +462,11 @@ func GetUserLogs(userId int, logType int, startTimestamp int64, endTimestamp int
return nil, 0, errors.New("查询日志失败")
}

logs, err = enrichTaskBillingLogs(logs)
if err != nil {
common.SysError("failed to enrich task billing logs: " + err.Error())
return nil, 0, err
}
formatUserLogs(logs, startIdx)
return logs, total, err
}


+ 50
- 0
model/log_identity_test.go Просмотреть файл

@@ -152,6 +152,56 @@ func TestGetUserLogsFiltersByChatIDAndUpstreamID(t *testing.T) {
require.Equal(t, "up-owner-a", logs[0].UpstreamId)
}

func TestGetUserLogsIncludesTaskSettlementOutsidePage(t *testing.T) {
db := setupLogIdentityDB(t)
preconsume := &Log{
UserId: 1, CreatedAt: 100, Type: LogTypeConsume, ModelName: "seedance",
Other: `{"billing_mode":"matrix","billing_phase":"preconsume","task_id":"task_cross_page"}`,
}
settlement := &Log{
UserId: 1, CreatedAt: 200, Type: LogTypeConsume, ModelName: "seedance",
Other: `{"billing_mode":"matrix","billing_phase":"settlement","task_id":"task_cross_page"}`,
}
require.NoError(t, db.Create(preconsume).Error)
require.NoError(t, db.Create(settlement).Error)

logs, total, err := GetUserLogs(1, LogTypeConsume, 0, 0, "", "", 1, 1, "", "", "", "")
require.NoError(t, err)
require.Equal(t, int64(2), total)
require.Len(t, logs, 2)
require.Equal(t, "preconsume", logBillingPhase(t, logs[0]))
require.Equal(t, "settlement", logBillingPhase(t, logs[1]))
}

func TestGetAllLogsIncludesTaskSettlementOutsidePage(t *testing.T) {
db := setupLogIdentityDB(t)
preconsume := &Log{
UserId: 1, CreatedAt: 100, Type: LogTypeConsume, ModelName: "seedance",
Other: `{"billing_mode":"matrix","billing_phase":"preconsume","task_id":"task_cross_page_admin"}`,
}
settlement := &Log{
UserId: 1, CreatedAt: 200, Type: LogTypeConsume, ModelName: "seedance",
Other: `{"billing_mode":"matrix","billing_phase":"settlement","task_id":"task_cross_page_admin"}`,
}
require.NoError(t, db.Create(preconsume).Error)
require.NoError(t, db.Create(settlement).Error)

logs, total, err := GetAllLogs(LogTypeConsume, 0, 0, "", "", "", 1, 1, 0, "", "", "", "")
require.NoError(t, err)
require.Equal(t, int64(2), total)
require.Len(t, logs, 2)
require.Equal(t, "preconsume", logBillingPhase(t, logs[0]))
require.Equal(t, "settlement", logBillingPhase(t, logs[1]))
}

func logBillingPhase(t *testing.T, log *Log) string {
t.Helper()
other := map[string]any{}
require.NoError(t, common.UnmarshalJsonStr(log.Other, &other))
phase, _ := other["billing_phase"].(string)
return phase
}

func TestRecordConsumeLogPersistsContextChatIDAndUpstreamID(t *testing.T) {
db := setupLogIdentityDB(t)



+ 73
- 3
model/main.go Просмотреть файл

@@ -207,7 +207,7 @@ func InitDB() (err error) {
return err
}
LoadEmailQuotaCache()
return nil
return nil
} else {
common.FatalLog(err)
}
@@ -259,13 +259,19 @@ func migrateDB() error {

err := DB.AutoMigrate(
&Channel{},
&ChannelAssetCredential{},
&Token{},
&User{},
&PasskeyCredential{},
&Option{},
&Redemption{},
&Ability{},
&Log{},
// &Log{} is intentionally omitted here. On PostgreSQL the logs table is a
// partitioned table managed by the pg_partman extension; GORM's AutoMigrate
// would try to alter its composite primary key (id, created_at) and create
// redundant per-partition indexes, breaking startup. The logs table is
// handled separately by ensureLogTable() below. On SQLite/MySQL it falls
// back to the original AutoMigrate behavior.
&Midjourney{},
&TopUp{},
&QuotaData{},
@@ -287,6 +293,7 @@ func migrateDB() error {
&EmailQuotaRule{},
&UserModelRateLimit{},
&UserAssetChannel{},
&UserAssetGroup{},
&UserMigrationBatch{},
&UserMigrationItem{},
&MigrationQuotaGrant{},
@@ -294,6 +301,9 @@ func migrateDB() error {
if err != nil {
return err
}
if err := ensureLogTable(); err != nil {
return err
}
if err := DB.Exec("DROP TABLE IF EXISTS channel_pricings").Error; err != nil {
return err
}
@@ -328,13 +338,15 @@ func migrateDBFast() error {
name string
}{
{&Channel{}, "Channel"},
{&ChannelAssetCredential{}, "ChannelAssetCredential"},
{&Token{}, "Token"},
{&User{}, "User"},
{&PasskeyCredential{}, "PasskeyCredential"},
{&Option{}, "Option"},
{&Redemption{}, "Redemption"},
{&Ability{}, "Ability"},
{&Log{}, "Log"},
// &Log{} omitted: see ensureLogTable() / migrateDB() for rationale
// (pg_partman-managed partitioned table on PostgreSQL).
{&Midjourney{}, "Midjourney"},
{&TopUp{}, "TopUp"},
{&QuotaData{}, "QuotaData"},
@@ -355,6 +367,7 @@ func migrateDBFast() error {
{&QuotaSyncLog{}, "QuotaSyncLog"},
{&EmailQuotaRule{}, "EmailQuotaRule"},
{&UserAssetChannel{}, "UserAssetChannel"},
{&UserAssetGroup{}, "UserAssetGroup"},
{&UserMigrationBatch{}, "UserMigrationBatch"},
{&UserMigrationItem{}, "UserMigrationItem"},
{&MigrationQuotaGrant{}, "MigrationQuotaGrant"},
@@ -382,6 +395,9 @@ func migrateDBFast() error {
return err
}
}
if err := ensureLogTable(); err != nil {
return err
}
if common.UsingSQLite {
if err := ensureSubscriptionPlanTableSQLite(); err != nil {
return err
@@ -395,7 +411,61 @@ func migrateDBFast() error {
return nil
}

// ensureLogTable handles the logs table, which behaves differently depending on
// the database backend:
//
// - SQLite / MySQL: behaves as before, GORM AutoMigrate creates/maintains it.
// - PostgreSQL: the logs table is a RANGE-partitioned table on created_at,
// created and maintained by the pg_partman extension (initdb script). GORM's
// AutoMigrate must NOT touch it, otherwise it would try to alter the composite
// primary key (id, created_at) - required by the partition key - and create
// redundant per-partition indexes from the model's index tags, breaking
// startup. Here we only verify the table exists and warn if it does not.
//
// No partition creation/maintenance logic lives in application code; pg_partman
// plus pg_cron own that responsibility entirely on the database side.
func ensureLogTable() error {
if !common.UsingPostgreSQL {
return DB.AutoMigrate(&Log{})
}
var exists bool
if err := DB.Raw(`SELECT EXISTS (
SELECT 1 FROM information_schema.tables
WHERE table_schema = current_schema() AND table_name = 'logs'
)`).Scan(&exists).Error; err != nil {
return err
}
if exists {
common.SysLog("logs table is a pg_partman-managed partitioned table, skipping GORM AutoMigrate")
return nil
}
common.SysLog("WARNING: PostgreSQL detected but 'logs' table not found. " +
"Ensure the pg_partman init script created the partitioned logs table before starting new-api.")
return nil
}

func migrateLOGDB() error {
// When LOG_SQL_DSN is empty, LOG_DB == DB and InitLogDB returns early without
// calling this function, so the logs table is already handled by ensureLogTable()
// during migrateDB(). This branch only runs for a dedicated PostgreSQL log
// database: keep pg_partman-managed behavior (skip AutoMigrate) consistent with
// the main database.
if common.LogSqlType == common.DatabaseTypePostgreSQL {
var exists bool
if err := LOG_DB.Raw(`SELECT EXISTS (
SELECT 1 FROM information_schema.tables
WHERE table_schema = current_schema() AND table_name = 'logs'
)`).Scan(&exists).Error; err != nil {
return err
}
if exists {
common.SysLog("logs table is a pg_partman-managed partitioned table, skipping GORM AutoMigrate")
return nil
}
common.SysLog("WARNING: PostgreSQL log database detected but 'logs' table not found. " +
"Ensure the pg_partman init script created the partitioned logs table before starting new-api.")
return nil
}
var err error
if err = LOG_DB.AutoMigrate(&Log{}); err != nil {
return err


+ 45
- 0
model/option_map_test.go Просмотреть файл

@@ -0,0 +1,45 @@
package model

import (
"testing"

"github.com/QuantumNous/new-api/common"
"github.com/QuantumNous/new-api/setting"
"github.com/QuantumNous/new-api/setting/operation_setting"
"github.com/stretchr/testify/require"
)

func TestUpdateOptionMapSynchronizesBooleanAndNumericRuntimeSettings(t *testing.T) {
oldTaskEnabled := common.TaskEnabled
oldPrice := operation_setting.Price
oldOptionMap := common.OptionMap
common.OptionMap = make(map[string]string)
t.Cleanup(func() {
common.TaskEnabled = oldTaskEnabled
operation_setting.Price = oldPrice
common.OptionMap = oldOptionMap
})

require.NoError(t, updateOptionMap("TaskEnabled", "false"))
require.False(t, common.TaskEnabled)
require.Equal(t, "false", common.OptionMap["TaskEnabled"])

require.NoError(t, updateOptionMap("Price", "2.5"))
require.Equal(t, 2.5, operation_setting.Price)
}

func TestUpdateOptionMapRejectsInvalidStructuredSettingsWithoutReplacingExistingValue(t *testing.T) {
oldOptionMap := common.OptionMap
common.OptionMap = make(map[string]string)
oldAutoGroups := setting.AutoGroups2JsonString()
t.Cleanup(func() {
common.OptionMap = oldOptionMap
require.NoError(t, setting.UpdateAutoGroupsByJsonString(oldAutoGroups))
})
require.NoError(t, setting.UpdateAutoGroupsByJsonString(`["default"]`))

err := updateOptionMap("AutoGroups", "not-json")

require.Error(t, err)
require.Equal(t, []string{"default"}, setting.GetAutoGroups())
}

+ 42
- 0
model/passkey_test.go Просмотреть файл

@@ -0,0 +1,42 @@
package model

import (
"testing"

"github.com/go-webauthn/webauthn/protocol"
webauthn "github.com/go-webauthn/webauthn/webauthn"
"github.com/stretchr/testify/require"
)

func TestPasskeyCredentialTransportsRoundTripAndIgnoreMalformedJSON(t *testing.T) {
credential := &PasskeyCredential{}
credential.SetTransports([]protocol.AuthenticatorTransport{protocol.USB, protocol.Internal})
require.Equal(t, []protocol.AuthenticatorTransport{protocol.USB, protocol.Internal}, credential.TransportList())

credential.Transports = "not-json"
require.Nil(t, credential.TransportList())

credential.SetTransports(nil)
require.Empty(t, credential.Transports)
}

func TestPasskeyCredentialConvertsWebAuthnFieldsWithoutLosingFlags(t *testing.T) {
webCredential := &webauthn.Credential{
ID: []byte("credential"), PublicKey: []byte("public-key"), AttestationType: "none",
Transport: []protocol.AuthenticatorTransport{protocol.Internal},
Flags: webauthn.CredentialFlags{UserPresent: true, UserVerified: true, BackupEligible: true},
Authenticator: webauthn.Authenticator{AAGUID: []byte("aaguid"), SignCount: 8, Attachment: protocol.Platform},
}

stored := NewPasskeyCredentialFromWebAuthn(9, webCredential)
require.NotNil(t, stored)
require.Equal(t, 9, stored.UserID)
require.True(t, stored.UserVerified)
require.EqualValues(t, 8, stored.SignCount)

roundTripped := stored.ToWebAuthnCredential()
require.Equal(t, webCredential.ID, roundTripped.ID)
require.Equal(t, webCredential.PublicKey, roundTripped.PublicKey)
require.True(t, roundTripped.Flags.UserVerified)
require.Equal(t, webCredential.Transport, roundTripped.Transport)
}

+ 41
- 0
model/prefill_group_test.go Просмотреть файл

@@ -0,0 +1,41 @@
package model

import (
"testing"

"github.com/stretchr/testify/require"
)

func TestJSONValueScanCopiesByteInputAndSupportsStringAndNil(t *testing.T) {
input := []byte(`{"items":["a"]}`)
var value JSONValue
require.NoError(t, value.Scan(input))
input[2] = 'X'
require.Equal(t, `{"items":["a"]}`, string(value))

require.NoError(t, value.Scan(`{"items":["b"]}`))
require.Equal(t, `{"items":["b"]}`, string(value))

require.NoError(t, value.Scan(nil))
require.Nil(t, value)
}

func TestJSONValueDatabaseAndJSONMarshallingPreservesRawPayload(t *testing.T) {
value := JSONValue(`{"models":["gpt-5"]}`)
databaseValue, err := value.Value()
require.NoError(t, err)
require.Equal(t, []byte(`{"models":["gpt-5"]}`), databaseValue)

encoded, err := value.MarshalJSON()
require.NoError(t, err)
require.Equal(t, []byte(`{"models":["gpt-5"]}`), encoded)

var decoded JSONValue
require.NoError(t, decoded.UnmarshalJSON([]byte(`["a","b"]`)))
require.Equal(t, `["a","b"]`, string(decoded))

var nilValue JSONValue
encoded, err = nilValue.MarshalJSON()
require.NoError(t, err)
require.Equal(t, []byte("null"), encoded)
}

+ 13
- 0
model/pricing_default_test.go Просмотреть файл

@@ -0,0 +1,13 @@
package model

import (
"testing"

"github.com/stretchr/testify/require"
)

func TestGetDefaultVendorIconReturnsKnownIconAndEmptyFallback(t *testing.T) {
require.Equal(t, "OpenAI", getDefaultVendorIcon("OpenAI"))
require.Equal(t, "Claude.Color", getDefaultVendorIcon("Anthropic"))
require.Equal(t, "", getDefaultVendorIcon("Unknown Vendor"))
}

+ 20
- 1
model/token.go Просмотреть файл

@@ -3,6 +3,7 @@ package model
import (
"errors"
"fmt"
"sort"
"strings"

"github.com/QuantumNous/new-api/common"
@@ -27,7 +28,7 @@ type Token struct {
AllowIps *string `json:"allow_ips" gorm:"default:''"`
UsedQuota int `json:"used_quota" gorm:"default:0"` // used quota
Group string `json:"group" gorm:"default:''"`
CrossGroupRetry bool `json:"cross_group_retry"` // 跨分组重试,仅auto分组有效
CrossGroupRetry bool `json:"cross_group_retry"` // 跨分组重试,仅auto分组有效
DeletedAt gorm.DeletedAt `gorm:"index"`
}

@@ -64,6 +65,24 @@ func GetAllUserTokens(userId int, startIdx int, num int) ([]*Token, error) {
return tokens, err
}

// GetUserConcreteTokenGroups returns the groups an administrator can bind.
// The auto group resolves at request time and cannot be bound directly.
func GetUserConcreteTokenGroups(userId int) ([]string, error) {
var rawGroups []string
if err := DB.Model(&Token{}).Where("user_id = ?", userId).Distinct().Pluck("group", &rawGroups).Error; err != nil {
return nil, err
}
groups := make([]string, 0, len(rawGroups))
for _, group := range rawGroups {
group = strings.TrimSpace(group)
if group != "" && group != "auto" {
groups = append(groups, group)
}
}
sort.Strings(groups)
return groups, nil
}

// sanitizeLikePattern 校验并清洗用户输入的 LIKE 搜索模式。
// 规则:
// 1. 转义 ! 和 _(使用 ! 作为 ESCAPE 字符,兼容 MySQL/PostgreSQL/SQLite)


+ 38
- 0
model/token_helpers_test.go Просмотреть файл

@@ -0,0 +1,38 @@
package model

import (
"testing"

"github.com/stretchr/testify/require"
)

func TestSanitizeLikePatternEscapesLiteralCharactersAndRejectsExpensiveWildcards(t *testing.T) {
pattern, err := sanitizeLikePattern("api_key!v1")
require.NoError(t, err)
require.Equal(t, "api!_key!!v1", pattern)

pattern, err = sanitizeLikePattern("ab%cd")
require.NoError(t, err)
require.Equal(t, "ab%cd", pattern)

_, err = sanitizeLikePattern("%%")
require.Error(t, err)
_, err = sanitizeLikePattern("a%b%c%d")
require.Error(t, err)
_, err = sanitizeLikePattern("a%")
require.Error(t, err)
}

func TestTokenParsersNormalizeIPAndModelLimitFields(t *testing.T) {
allowIPs := " 127.0.0.1,\n 10.0.0.1 \n\n"
token := Token{AllowIps: &allowIPs, ModelLimits: "gpt-5,claude-sonnet"}

require.Equal(t, []string{"127.0.0.1", "10.0.0.1"}, token.GetIpLimits())
require.Equal(t, []string{"gpt-5", "claude-sonnet"}, token.GetModelLimits())
require.Equal(t, map[string]bool{"gpt-5": true, "claude-sonnet": true}, token.GetModelLimitsMap())

token.AllowIps = nil
token.ModelLimits = ""
require.Empty(t, token.GetIpLimits())
require.Empty(t, token.GetModelLimits())
}

+ 28
- 11
model/topup.go Просмотреть файл

@@ -13,18 +13,27 @@ import (
)

type TopUp struct {
Id int `json:"id"`
UserId int `json:"user_id" gorm:"index"`
Amount int64 `json:"amount"`
Money float64 `json:"money"`
TradeNo string `json:"trade_no" gorm:"unique;type:varchar(255);index"`
PaymentMethod string `json:"payment_method" gorm:"type:varchar(50)"`
CreateTime int64 `json:"create_time"`
CompleteTime int64 `json:"complete_time"`
Status string `json:"status"`
UserEmail string `json:"user_email" gorm:"-"` // Join 查询时填充,非数据库字段
Id int `json:"id"`
UserId int `json:"user_id" gorm:"index"`
Amount int64 `json:"amount"`
Money float64 `json:"money"`
TradeNo string `json:"trade_no" gorm:"unique;type:varchar(255);index"`
PaymentMethod string `json:"payment_method" gorm:"type:varchar(50)"`
PaymentProvider string `json:"payment_provider" gorm:"type:varchar(50);default:''"`
CreateTime int64 `json:"create_time"`
CompleteTime int64 `json:"complete_time"`
Status string `json:"status"`
UserEmail string `json:"user_email" gorm:"-"` // Join 查询时填充,非数据库字段
}

const (
PaymentProviderEpay = "epay"
PaymentProviderStripe = "stripe"
PaymentProviderCreem = "creem"
PaymentProviderWechat = "wechat_pay"
PaymentProviderAlipay = "alipay"
)

// fillTopUpEmails 批量填充 topup 记录的用户邮箱
func fillTopUpEmails(topups []*TopUp) {
if len(topups) == 0 {
@@ -113,7 +122,10 @@ func Recharge(referenceId string, customerId string) (err error) {
}

quota = topUp.Money * common.QuotaPerUnit
err = tx.Model(&User{}).Where("id = ?", topUp.UserId).Updates(map[string]interface{}{"stripe_customer": customerId, "quota": gorm.Expr("quota + ?", quota)}).Error
if topUp.PaymentProvider != "" && topUp.PaymentProvider != PaymentProviderStripe {
return fmt.Errorf("支付网关不匹配: 订单由 %s 创建, 但 Stripe webhook 尝试完成", topUp.PaymentProvider)
}
err = tx.Model(&User{}).Where("id = ?", topUp.UserId).Updates(map[string]interface{}{"stripe_customer": customerId, "quota": gorm.Expr("quota + ?", quota)}).Error
if err != nil {
return err
}
@@ -360,6 +372,11 @@ func RechargeCreem(referenceId string, customerEmail string, customerName string
return errors.New("充值订单状态错误")
}

// 防止跨网关回调攻击
if topUp.PaymentProvider != "" && topUp.PaymentProvider != PaymentProviderCreem {
return fmt.Errorf("支付网关不匹配: 订单由 %s 创建, 但 Creem webhook 尝试完成", topUp.PaymentProvider)
}

topUp.CompleteTime = common.GetTimestamp()
topUp.Status = common.TopUpStatusSuccess
err = tx.Save(topUp).Error


+ 5
- 1
model/topup_wechat.go Просмотреть файл

@@ -59,7 +59,11 @@ func rechargeByQRCodePayment(tradeNo string, paymentMethod string) error {

dMoney := decimal.NewFromFloat(topUp.Money)
dQuotaPerUnit := decimal.NewFromFloat(common.QuotaPerUnit)
quotaToAdd = dMoney.Mul(dQuotaPerUnit).IntPart()
// 防止跨网关回调攻击:扫码支付 webhook 只能完成对应渠道创建的订单
if topUp.PaymentProvider != "" && topUp.PaymentProvider != PaymentProviderWechat && topUp.PaymentProvider != PaymentProviderAlipay {
return fmt.Errorf("支付网关不匹配: 订单由 %s 创建, 但 %s webhook 尝试完成", topUp.PaymentProvider, paymentMethod)
}
quotaToAdd = dMoney.Mul(dQuotaPerUnit).IntPart()

if quotaToAdd <= 0 {
return errors.New("无效的充值额度")


+ 21
- 0
model/twofa_test.go Просмотреть файл

@@ -0,0 +1,21 @@
package model

import (
"testing"
"time"

"github.com/stretchr/testify/require"
)

func TestTwoFAIsLockedDependsOnFutureLockDeadline(t *testing.T) {
noDeadline := &TwoFA{}
require.False(t, noDeadline.IsLocked())

future := time.Now().Add(time.Minute)
locked := &TwoFA{LockedUntil: &future}
require.True(t, locked.IsLocked())

past := time.Now().Add(-time.Minute)
expired := &TwoFA{LockedUntil: &past}
require.False(t, expired.IsLocked())
}

+ 8
- 0
model/user_asset_channel.go Просмотреть файл

@@ -73,6 +73,14 @@ func GetUserAssetChannelsByTypes(userId int, channelTypes []int, group string) (
return bindings, err
}

func GetUserAssetChannels(userId int, group string) ([]UserAssetChannel, error) {
var bindings []UserAssetChannel
err := DB.Where("user_id = ? AND "+commonGroupCol+" = ?", userId, group).
Order("updated_at DESC, id DESC").
Find(&bindings).Error
return bindings, err
}

func DeleteUserAssetChannelsByTypesWithTx(tx *gorm.DB, userId int, channelTypes []int, group string) error {
if len(channelTypes) == 0 {
return nil


+ 20
- 0
model/user_asset_channel_test.go Просмотреть файл

@@ -92,6 +92,26 @@ func TestGetUserAssetChannelsByTypesSortsLatestFirst(t *testing.T) {
assert.Equal(t, 2, bindings[1].ChannelType)
}

func TestGetUserAssetChannelsReturnsAllTypesSortedLatestFirst(t *testing.T) {
db := setupUserAssetChannelDB(t)
require.NoError(t, BindUserAssetChannel(1, 2, "default", 100))
require.NoError(t, BindUserAssetChannel(1, 3, "default", 200))
require.NoError(t, BindUserAssetChannel(1, 4, "vip", 300))
require.NoError(t, db.Model(&UserAssetChannel{}).
Where("user_id = ? AND channel_type = ?", 1, 2).
Update("updated_at", int64(100)).Error)
require.NoError(t, db.Model(&UserAssetChannel{}).
Where("user_id = ? AND channel_type = ?", 1, 3).
Update("updated_at", int64(200)).Error)

bindings, err := GetUserAssetChannels(1, "default")

require.NoError(t, err)
require.Len(t, bindings, 2)
assert.Equal(t, 3, bindings[0].ChannelType)
assert.Equal(t, 2, bindings[1].ChannelType)
}

func TestBindUserAssetChannelWithTxUpserts(t *testing.T) {
db := setupUserAssetChannelDB(t)



+ 45
- 0
model/user_asset_group.go Просмотреть файл

@@ -0,0 +1,45 @@
package model

import (
"errors"

"github.com/QuantumNous/new-api/common"
"gorm.io/gorm"
)

// UserAssetGroup binds a platform user to an upstream asset group per channel.
type UserAssetGroup struct {
Id int `json:"id" gorm:"primaryKey"`
UserId int `json:"user_id" gorm:"not null;uniqueIndex:idx_user_asset_group,priority:1"`
ChannelId int `json:"channel_id" gorm:"not null;uniqueIndex:idx_user_asset_group,priority:2"`
GroupId string `json:"group_id" gorm:"type:varchar(255);not null"`
CreatedAt int64 `json:"created_at"`
UpdatedAt int64 `json:"updated_at"`
}

func (UserAssetGroup) TableName() string {
return "user_asset_groups"
}

func GetUserAssetGroup(userId int, channelId int) (*UserAssetGroup, error) {
var binding UserAssetGroup
err := DB.Where("user_id = ? AND channel_id = ?", userId, channelId).First(&binding).Error
if errors.Is(err, gorm.ErrRecordNotFound) {
return nil, nil
}
if err != nil {
return nil, err
}
return &binding, nil
}

func CreateUserAssetGroup(userId int, channelId int, groupId string) error {
now := common.GetTimestamp()
return DB.Create(&UserAssetGroup{
UserId: userId,
ChannelId: channelId,
GroupId: groupId,
CreatedAt: now,
UpdatedAt: now,
}).Error
}

+ 63
- 0
model/user_asset_group_test.go Просмотреть файл

@@ -0,0 +1,63 @@
package model

import (
"testing"

"github.com/glebarez/sqlite"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"gorm.io/gorm"
"gorm.io/gorm/logger"
)

func setupUserAssetGroupDB(t *testing.T) *gorm.DB {
t.Helper()

db, err := gorm.Open(sqlite.Open("file:user_asset_groups?mode=memory&cache=shared"), &gorm.Config{
Logger: logger.Default.LogMode(logger.Silent),
})
require.NoError(t, err)

sqlDB, err := db.DB()
require.NoError(t, err)
sqlDB.SetMaxOpenConns(1)

origDB := DB
DB = db
require.NoError(t, db.AutoMigrate(&UserAssetGroup{}))

t.Cleanup(func() {
DB = origDB
require.NoError(t, sqlDB.Close())
})

return db
}

func TestUserAssetGroupScopesByUserAndChannel(t *testing.T) {
setupUserAssetGroupDB(t)

require.NoError(t, CreateUserAssetGroup(10, 101, "group-a"))
require.NoError(t, CreateUserAssetGroup(11, 101, "group-b"))
require.NoError(t, CreateUserAssetGroup(10, 102, "group-c"))

binding, err := GetUserAssetGroup(10, 101)
require.NoError(t, err)
require.NotNil(t, binding)
assert.Equal(t, "group-a", binding.GroupId)
}

func TestUserAssetGroupRejectsDuplicateUserChannel(t *testing.T) {
setupUserAssetGroupDB(t)

require.NoError(t, CreateUserAssetGroup(10, 101, "first"))
assert.Error(t, CreateUserAssetGroup(10, 101, "second"))
}

func TestGetUserAssetGroupReturnsNilWhenNotFound(t *testing.T) {
setupUserAssetGroupDB(t)

binding, err := GetUserAssetGroup(10, 101)
require.NoError(t, err)
assert.Nil(t, binding)
}

+ 42
- 0
model/user_cache_test.go Просмотреть файл

@@ -0,0 +1,42 @@
package model

import (
"net/http/httptest"
"testing"

"github.com/QuantumNous/new-api/common"
"github.com/QuantumNous/new-api/constant"
"github.com/gin-gonic/gin"
"github.com/stretchr/testify/require"
)

func TestUserBaseGetSettingHandlesValidAndInvalidJSON(t *testing.T) {
user := UserBase{Setting: `{"language":"en","billing_preference":"subscription"}`}
setting := user.GetSetting()
require.Equal(t, "en", setting.Language)
require.Equal(t, "subscription", setting.BillingPreference)

invalid := UserBase{Setting: "{"}
require.Equal(t, "", invalid.GetSetting().Language)
}

func TestUserBaseWriteContextPopulatesRelayFields(t *testing.T) {
context, _ := gin.CreateTestContext(httptest.NewRecorder())
user := UserBase{
Source: "oauth", Group: "vip", Quota: 123, Status: common.UserStatusEnabled,
Email: "user@example.com", Username: "user", Setting: `{"language":"en"}`,
}

user.WriteContext(context)

require.Equal(t, "vip", common.GetContextKeyString(context, constant.ContextKeyUserGroup))
require.Equal(t, 123, common.GetContextKeyInt(context, constant.ContextKeyUserQuota))
require.Equal(t, common.UserStatusEnabled, common.GetContextKeyInt(context, constant.ContextKeyUserStatus))
require.Equal(t, "user@example.com", common.GetContextKeyString(context, constant.ContextKeyUserEmail))
require.Equal(t, "user", common.GetContextKeyString(context, constant.ContextKeyUserName))
require.Equal(t, "oauth", common.GetContextKeyString(context, constant.ContextKeyUserSource))
}

func TestGetUserCacheKeyIsNamespacedByUserID(t *testing.T) {
require.Equal(t, "user:42", getUserCacheKey(42))
}

+ 38
- 0
model/utils_test.go Просмотреть файл

@@ -0,0 +1,38 @@
package model

import (
"errors"
"testing"

"github.com/QuantumNous/new-api/common"
"github.com/stretchr/testify/require"
"gorm.io/gorm"
)

func TestRecordExistClassifiesSuccessNotFoundAndDatabaseError(t *testing.T) {
exists, err := RecordExist(nil)
require.True(t, exists)
require.NoError(t, err)

exists, err = RecordExist(gorm.ErrRecordNotFound)
require.False(t, exists)
require.NoError(t, err)

dbErr := errors.New("database unavailable")
exists, err = RecordExist(dbErr)
require.False(t, exists)
require.ErrorIs(t, err, dbErr)
}

func TestShouldUpdateRedisRequiresRedisDatabaseSourceAndNoError(t *testing.T) {
oldRedisEnabled := common.RedisEnabled
t.Cleanup(func() { common.RedisEnabled = oldRedisEnabled })

common.RedisEnabled = false
require.False(t, shouldUpdateRedis(true, nil))

common.RedisEnabled = true
require.False(t, shouldUpdateRedis(false, nil))
require.False(t, shouldUpdateRedis(true, errors.New("query failed")))
require.True(t, shouldUpdateRedis(true, nil))
}

Двоичные данные
new-api.exe~ Просмотреть файл


+ 494
- 0
relay/channel/task/chinamobile_seedance/adaptor.go Просмотреть файл

@@ -0,0 +1,494 @@
package chinamobile_seedance

import (
"bytes"
"fmt"
"io"
"net/http"
"strconv"
"strings"
"time"

"github.com/QuantumNous/new-api/common"
"github.com/QuantumNous/new-api/constant"
"github.com/QuantumNous/new-api/dto"
"github.com/QuantumNous/new-api/model"
taskcommon "github.com/QuantumNous/new-api/relay/channel/task/taskcommon"
relaycommon "github.com/QuantumNous/new-api/relay/common"
"github.com/QuantumNous/new-api/service"
relaytypes "github.com/QuantumNous/new-api/types"
"github.com/gin-gonic/gin"
"github.com/pkg/errors"
)

type responsePayload struct {
ID string `json:"id"`
}

type upstreamError struct {
Code string `json:"code"`
Message string `json:"message"`
}

func (e *upstreamError) UnmarshalJSON(data []byte) error {
if len(bytes.TrimSpace(data)) == 0 || string(bytes.TrimSpace(data)) == "null" {
return nil
}
var message string
if err := common.Unmarshal(data, &message); err == nil {
e.Message = message
return nil
}
type alias upstreamError
var parsed alias
if err := common.Unmarshal(data, &parsed); err != nil {
return err
}
*e = upstreamError(parsed)
return nil
}

type responseTask struct {
ID string `json:"id"`
Model string `json:"model"`
Status string `json:"status"`
Content struct {
VideoURL string `json:"video_url"`
} `json:"content"`
Usage struct {
CompletionTokens int `json:"completion_tokens"`
TotalTokens int `json:"total_tokens"`
} `json:"usage"`
Error upstreamError `json:"error"`
Message string `json:"message"`
CreatedAt int64 `json:"created_at"`
UpdatedAt int64 `json:"updated_at"`
}

type TaskAdaptor struct {
taskcommon.BaseBilling
ChannelType int
apiKey string
baseURL string
newSDK seedanceSDKFactory
submitTimeout time.Duration
queryTimeout time.Duration
}

func (a *TaskAdaptor) Init(info *relaycommon.RelayInfo) {
if info != nil && info.ChannelMeta != nil {
a.ChannelType = info.ChannelType
a.baseURL = strings.TrimRight(info.ChannelBaseUrl, "/")
a.apiKey = info.ApiKey
}
if a.newSDK == nil {
a.newSDK = newRealSDK
}
if a.submitTimeout == 0 {
a.submitTimeout = defaultSubmitTimeout
}
if a.queryTimeout == 0 {
a.queryTimeout = defaultQueryTimeout
}
}

func (a *TaskAdaptor) ValidateRequestAndSetAction(c *gin.Context, info *relaycommon.RelayInfo) (taskErr *dto.TaskError) {
if _, err := relaycommon.GetTaskRequest(c); err == nil {
info.Action = constant.TaskActionGenerate
return nil
}
return relaycommon.ValidateBasicTaskRequest(c, info, constant.TaskActionGenerate)
}

func (a *TaskAdaptor) BuildRequestURL(info *relaycommon.RelayInfo) (string, error) {
baseURL := a.baseURL
if strings.TrimSpace(baseURL) == "" && info != nil {
baseURL = strings.TrimRight(info.ChannelBaseUrl, "/")
}
if strings.TrimSpace(baseURL) == "" {
baseURL = defaultBaseURL
}
return fmt.Sprintf("%s/contents/generations/tasks", strings.TrimRight(baseURL, "/")), nil
}

func (a *TaskAdaptor) BuildRequestHeader(c *gin.Context, req *http.Request, info *relaycommon.RelayInfo) error {
req.Header.Set("Content-Type", "application/json")
req.Header.Set("Accept", "application/json")
return nil
}

func (a *TaskAdaptor) BuildRequestBody(c *gin.Context, info *relaycommon.RelayInfo) (io.Reader, error) {
req, err := relaycommon.GetTaskRequest(c)
if err != nil {
return nil, err
}
body, err := a.convertToRequestPayload(&req)
if err != nil {
return nil, errors.Wrap(err, "convert request payload failed")
}
if info != nil {
if info.IsModelMapped {
body["model"] = info.UpstreamModelName
} else if modelName, _ := body["model"].(string); modelName != "" {
info.UpstreamModelName = modelName
}
}
data, err := common.Marshal(body)
if err != nil {
return nil, err
}
return bytes.NewReader(data), nil
}

func (a *TaskAdaptor) DoRequest(c *gin.Context, info *relaycommon.RelayInfo, requestBody io.Reader) (*http.Response, error) {
data, err := io.ReadAll(requestBody)
if err != nil {
return nil, err
}
payload := map[string]interface{}{}
if err := common.Unmarshal(data, &payload); err != nil {
return nil, err
}
modelName := chinaMobileSeedanceModel(modelFromPayloadOrInfo(payload, info))
payload["model"] = modelName
if info != nil {
info.UpstreamModelName = modelName
}
client, err := a.createSDK(info, modelName)
if err != nil {
return nil, err
}
submitTimeout := a.submitTimeout
if submitTimeout == 0 {
submitTimeout = defaultSubmitTimeout
}
taskID, err := runSDKCall[string](submitTimeout, func() (string, error) {
return client.CreateVideoGenerationTask(payload)
})
if err != nil {
return chinaMobileSeedanceErrorResponse(err), nil
}
respBody, err := common.Marshal(map[string]any{"id": taskID})
if err != nil {
return nil, err
}
return &http.Response{
StatusCode: http.StatusOK,
Header: make(http.Header),
Body: io.NopCloser(bytes.NewReader(respBody)),
}, nil
}

func (a *TaskAdaptor) DoResponse(c *gin.Context, resp *http.Response, info *relaycommon.RelayInfo) (taskID string, taskData []byte, taskErr *dto.TaskError) {
responseBody, err := io.ReadAll(resp.Body)
if err != nil {
return "", nil, service.TaskErrorWrapper(err, "read_response_body_failed", http.StatusInternalServerError)
}
_ = resp.Body.Close()

var dResp responsePayload
if err := common.Unmarshal(responseBody, &dResp); err != nil {
return "", nil, service.TaskErrorWrapper(errors.Wrapf(err, "body: %s", responseBody), "unmarshal_response_body_failed", http.StatusInternalServerError)
}
if strings.TrimSpace(dResp.ID) == "" {
return "", nil, service.TaskErrorWrapper(fmt.Errorf("task_id is empty"), "invalid_response", http.StatusInternalServerError)
}

clientPayload := map[string]any{}
if err := common.Unmarshal(responseBody, &clientPayload); err != nil {
return "", nil, service.TaskErrorWrapper(err, "unmarshal_response_body_failed", http.StatusInternalServerError)
}
clientPayload = relaytypes.CloneMapAny(clientPayload)
if info.PublicTaskID != "" {
clientPayload["id"] = info.PublicTaskID
}
if _, ok := clientPayload["created_at"]; !ok {
clientPayload["created_at"] = time.Now().Unix()
}
if _, ok := clientPayload["model"]; !ok {
clientPayload["model"] = info.OriginModelName
}
c.JSON(http.StatusOK, clientPayload)
return dResp.ID, responseBody, nil
}

func (a *TaskAdaptor) FetchTask(baseUrl, key string, body map[string]any, proxy string) (*http.Response, error) {
taskID, _ := body["task_id"].(string)
if strings.TrimSpace(taskID) == "" {
return nil, fmt.Errorf("invalid task_id")
}
modelName, _ := body["model"].(string)
client, err := a.createSDKFromValues(baseUrl, key, modelName)
if err != nil {
return nil, err
}
queryTimeout := a.queryTimeout
if queryTimeout == 0 {
queryTimeout = defaultQueryTimeout
}
result, err := runSDKCall[map[string]interface{}](queryTimeout, func() (map[string]interface{}, error) {
return client.QueryVideoGenerationTask(taskID)
})
if err != nil {
return nil, err
}
respBody, err := common.Marshal(result)
if err != nil {
return nil, err
}
return &http.Response{
StatusCode: http.StatusOK,
Header: make(http.Header),
Body: io.NopCloser(bytes.NewReader(respBody)),
}, nil
}

func (a *TaskAdaptor) ParseTaskResult(respBody []byte) (*relaycommon.TaskInfo, error) {
resTask := responseTask{}
if err := common.Unmarshal(respBody, &resTask); err != nil {
return nil, errors.Wrap(err, "unmarshal task result failed")
}

taskResult := relaycommon.TaskInfo{Code: 0}
switch strings.ToLower(resTask.Status) {
case "pending", "queued":
taskResult.Status = model.TaskStatusQueued
taskResult.Progress = "10%"
case "processing", "running":
taskResult.Status = model.TaskStatusInProgress
taskResult.Progress = "50%"
case "succeeded", "success":
taskResult.Status = model.TaskStatusSuccess
taskResult.Progress = "100%"
taskResult.Url = resTask.Content.VideoURL
taskResult.CompletionTokens = resTask.Usage.CompletionTokens
taskResult.TotalTokens = resTask.Usage.TotalTokens
case "failed", "expired", "cancelled":
taskResult.Status = model.TaskStatusFailure
taskResult.Progress = "100%"
taskResult.Reason = upstreamReason(resTask)
default:
if resTask.Error.Message != "" || resTask.Message != "" {
taskResult.Status = model.TaskStatusFailure
taskResult.Progress = "100%"
taskResult.Reason = upstreamReason(resTask)
} else {
taskResult.Status = model.TaskStatusInProgress
taskResult.Progress = "30%"
}
}
return &taskResult, nil
}

func (a *TaskAdaptor) ConvertToOpenAIVideo(originTask *model.Task) ([]byte, error) {
var dResp responseTask
if err := common.Unmarshal(originTask.Data, &dResp); err != nil {
return nil, errors.Wrap(err, "unmarshal chinamobile seedance task data failed")
}

openAIVideo := dto.NewOpenAIVideo()
openAIVideo.ID = originTask.TaskID
openAIVideo.TaskID = originTask.TaskID
openAIVideo.Status = originTask.Status.ToVideoStatus()
openAIVideo.SetProgressStr(originTask.Progress)
openAIVideo.SetMetadata("url", dResp.Content.VideoURL)
openAIVideo.CreatedAt = originTask.CreatedAt
openAIVideo.CompletedAt = originTask.UpdatedAt
openAIVideo.Model = originTask.Properties.OriginModelName

if originTask.Status == model.TaskStatusFailure || dResp.Status == "failed" {
message := upstreamReason(dResp)
if message == "" {
message = "task failed"
}
code := dResp.Error.Code
if code == "" {
code = "failed"
}
openAIVideo.Error = &dto.OpenAIVideoError{Message: message, Code: code}
}
return common.Marshal(openAIVideo)
}

func (a *TaskAdaptor) GetModelList() []string {
return ModelList
}

func (a *TaskAdaptor) GetChannelName() string {
return ChannelName
}

func (a *TaskAdaptor) createSDK(info *relaycommon.RelayInfo, modelName string) (seedanceSDK, error) {
baseURL := a.baseURL
apiKey := a.apiKey
if info != nil && info.ChannelMeta != nil {
baseURL = info.ChannelBaseUrl
apiKey = info.ApiKey
}
return a.createSDKFromValues(baseURL, apiKey, modelName)
}

func (a *TaskAdaptor) createSDKFromValues(baseURL, apiKey, modelName string) (seedanceSDK, error) {
factory := a.newSDK
if factory == nil {
factory = newRealSDK
}
return factory(strings.TrimRight(baseURL, "/"), apiKey, modelName)
}

func (a *TaskAdaptor) convertToRequestPayload(req *relaycommon.TaskSubmitReq) (map[string]interface{}, error) {
payload := map[string]interface{}{
"model": req.Model,
}
for k, v := range req.Metadata {
payload[k] = v
}
if sec, _ := strconv.Atoi(req.Seconds); sec > 0 {
if _, ok := payload["duration"]; !ok {
payload["duration"] = sec
}
}
content, err := normalizeContent(payload["content"])
if err != nil {
return nil, err
}
if len(content) == 0 && req.HasImage() {
for _, imgURL := range req.Images {
content = append(content, map[string]interface{}{
"type": "image_url",
"image_url": map[string]interface{}{"url": imgURL},
})
}
}
if strings.TrimSpace(req.Prompt) != "" && !hasTextContent(content) {
content = append(content, map[string]interface{}{
"type": "text",
"text": req.Prompt,
})
}
if len(content) > 0 {
payload["content"] = content
}
return payload, nil
}

func normalizeContent(input any) ([]interface{}, error) {
if input == nil {
return nil, nil
}
data, err := common.Marshal(input)
if err != nil {
return nil, err
}
var content []interface{}
if err := common.Unmarshal(data, &content); err != nil {
return nil, err
}
return content, nil
}

func hasTextContent(content []interface{}) bool {
for _, item := range content {
m, ok := item.(map[string]interface{})
if !ok {
continue
}
if m["type"] == "text" && strings.TrimSpace(fmt.Sprint(m["text"])) != "" {
return true
}
}
return false
}

func modelFromPayloadOrInfo(payload map[string]interface{}, info *relaycommon.RelayInfo) string {
if modelName, _ := payload["model"].(string); strings.TrimSpace(modelName) != "" {
return modelName
}
if info != nil {
if strings.TrimSpace(info.UpstreamModelName) != "" {
return info.UpstreamModelName
}
return info.OriginModelName
}
return defaultModel
}

func chinaMobileSeedanceModel(modelName string) string {
switch strings.TrimSpace(modelName) {
case "", "doubao-seedance-2-0-260128", "doubao-seedance-2-0-fast-260128":
return defaultModel
default:
return modelName
}
}

type chinaMobileSeedanceUpstreamError struct {
ErrorCode string `json:"ErrorCode"`
ErrorMessage string `json:"ErrorMessage"`
}

func chinaMobileSeedanceErrorResponse(err error) *http.Response {
statusCode := http.StatusBadGateway
body := map[string]any{
"code": "upstream_error",
"message": err.Error(),
}
if upstreamErr, ok := parseChinaMobileSeedanceError(err); ok {
statusCode = chinaMobileSeedanceStatusCode(upstreamErr.ErrorCode)
body["code"] = upstreamErr.ErrorCode
body["message"] = upstreamErr.ErrorMessage
}
respBody, marshalErr := common.Marshal(body)
if marshalErr != nil {
respBody = []byte(`{"code":"upstream_error","message":"upstream error"}`)
}
return &http.Response{
StatusCode: statusCode,
Header: make(http.Header),
Body: io.NopCloser(bytes.NewReader(respBody)),
}
}

func parseChinaMobileSeedanceError(err error) (chinaMobileSeedanceUpstreamError, bool) {
if err == nil {
return chinaMobileSeedanceUpstreamError{}, false
}
text := err.Error()
start := strings.Index(text, "{")
end := strings.LastIndex(text, "}")
if start < 0 || end < start {
return chinaMobileSeedanceUpstreamError{}, false
}
var upstreamErr chinaMobileSeedanceUpstreamError
if err := common.Unmarshal([]byte(text[start:end+1]), &upstreamErr); err != nil {
return chinaMobileSeedanceUpstreamError{}, false
}
return upstreamErr, upstreamErr.ErrorCode != "" || upstreamErr.ErrorMessage != ""
}

func chinaMobileSeedanceStatusCode(code string) int {
upperCode := strings.ToUpper(strings.TrimSpace(code))
switch {
case strings.Contains(upperCode, "PERMISSION") || strings.Contains(upperCode, "UNAUTHORIZED"):
return http.StatusForbidden
case strings.Contains(upperCode, "SENSITIVE") || strings.Contains(upperCode, "INVALID") || strings.Contains(upperCode, "BAD"):
return http.StatusBadRequest
case strings.Contains(upperCode, "RATE") || strings.Contains(upperCode, "LIMIT"):
return http.StatusTooManyRequests
default:
return http.StatusBadGateway
}
}

func upstreamReason(resTask responseTask) string {
if resTask.Error.Message != "" {
return resTask.Error.Message
}
if resTask.Message != "" {
return resTask.Message
}
if resTask.Status != "" {
return "task " + resTask.Status
}
return "upstream error"
}

+ 349
- 0
relay/channel/task/chinamobile_seedance/adaptor_test.go Просмотреть файл

@@ -0,0 +1,349 @@
package chinamobile_seedance

import (
"errors"
"io"
"net/http"
"net/http/httptest"
"strings"
"testing"
"time"

"github.com/QuantumNous/new-api/common"
"github.com/QuantumNous/new-api/constant"
"github.com/QuantumNous/new-api/model"
relaycommon "github.com/QuantumNous/new-api/relay/common"
"github.com/gin-gonic/gin"
"github.com/stretchr/testify/require"
)

func TestRunSDKCallRecoversPanic(t *testing.T) {
_, err := runSDKCall[string](time.Second, func() (string, error) {
panic("attestation failed")
})
require.Error(t, err)
require.Contains(t, err.Error(), "upstream_sdk_panic")
}

func TestRunSDKCallTimesOut(t *testing.T) {
start := time.Now()
_, err := runSDKCall[string](20*time.Millisecond, func() (string, error) {
time.Sleep(200 * time.Millisecond)
return "late", nil
})
require.Error(t, err)
require.Contains(t, err.Error(), "upstream_timeout")
require.Less(t, time.Since(start), 150*time.Millisecond)
}

func TestRunSDKCallTimeoutDoesNotPermanentlyOccupySemaphore(t *testing.T) {
original := sdkSemaphore
sdkSemaphore = make(chan struct{}, 1)
t.Cleanup(func() {
sdkSemaphore = original
})

block := make(chan struct{})
_, err := runSDKCall[string](20*time.Millisecond, func() (string, error) {
<-block
return "late", nil
})
require.ErrorContains(t, err, "upstream_timeout")

_, err = runSDKCall[string](time.Second, func() (string, error) {
return "ok", nil
})
require.NoError(t, err)

close(block)
}

func TestRunSDKCallReturnsError(t *testing.T) {
_, err := runSDKCall[string](time.Second, func() (string, error) {
return "", errors.New("upstream rejected")
})
require.ErrorContains(t, err, "upstream rejected")
}

type fakeSDK struct {
createInput map[string]interface{}
createID string
queryTaskID string
queryResult map[string]interface{}
err error
}

func (f *fakeSDK) CreateVideoGenerationTask(data map[string]interface{}) (string, error) {
f.createInput = data
if f.err != nil {
return "", f.err
}
return f.createID, nil
}

func (f *fakeSDK) QueryVideoGenerationTask(taskID string) (map[string]interface{}, error) {
f.queryTaskID = taskID
if f.err != nil {
return nil, f.err
}
return f.queryResult, nil
}

func TestBuildRequestBodyNormalizesContentToInterfaceSlice(t *testing.T) {
adaptor := &TaskAdaptor{}
adaptor.Init(&relaycommon.RelayInfo{
ChannelMeta: &relaycommon.ChannelMeta{ChannelType: constant.ChannelTypeChinaMobileSeedance},
})
c := newTaskRequestContext(t, `{
"model":"doubao-seedance-2.0",
"prompt":"current prompt",
"metadata":{
"content":[
{"type":"video_url","video_url":{"url":"https://example.test/input.mp4"},"role":"reference_video"}
],
"duration":5
}
}`)
info := &relaycommon.RelayInfo{
ChannelMeta: &relaycommon.ChannelMeta{UpstreamModelName: "doubao-seedance-2.0"},
}

body, err := adaptor.BuildRequestBody(c, info)
require.NoError(t, err)
data, err := io.ReadAll(body)
require.NoError(t, err)

var payload map[string]interface{}
require.NoError(t, common.Unmarshal(data, &payload))
content, ok := payload["content"].([]interface{})
require.True(t, ok)
require.Len(t, content, 2)
require.Equal(t, "video_url", content[0].(map[string]interface{})["type"])
require.Equal(t, "text", content[1].(map[string]interface{})["type"])
require.Equal(t, "current prompt", content[1].(map[string]interface{})["text"])
require.Equal(t, float64(5), payload["duration"])
}

func TestDoRequestCallsSDKAndReturnsSyntheticResponse(t *testing.T) {
fake := &fakeSDK{createID: "upstream-task-1"}
adaptor := &TaskAdaptor{
newSDK: func(baseURL, apiKey, model string) (seedanceSDK, error) {
require.Equal(t, "https://cm.example.com/api/v3", baseURL)
require.Equal(t, "sk-test", apiKey)
require.Equal(t, "doubao-seedance-2.0", model)
return fake, nil
},
}
adaptor.Init(&relaycommon.RelayInfo{
ChannelMeta: &relaycommon.ChannelMeta{
ChannelType: constant.ChannelTypeChinaMobileSeedance,
ChannelBaseUrl: "https://cm.example.com/api/v3",
ApiKey: "sk-test",
},
OriginModelName: "doubao-seedance-2.0",
})
payload := strings.NewReader(`{"model":"doubao-seedance-2.0","content":[{"type":"text","text":"hello"}]}`)

resp, err := adaptor.DoRequest(&gin.Context{}, &relaycommon.RelayInfo{
ChannelMeta: &relaycommon.ChannelMeta{
ChannelType: constant.ChannelTypeChinaMobileSeedance,
ChannelBaseUrl: "https://cm.example.com/api/v3",
ApiKey: "sk-test",
},
OriginModelName: "doubao-seedance-2.0",
}, payload)

require.NoError(t, err)
require.Equal(t, http.StatusOK, resp.StatusCode)
data, err := io.ReadAll(resp.Body)
require.NoError(t, err)
require.JSONEq(t, `{"id":"upstream-task-1"}`, string(data))
require.Equal(t, "doubao-seedance-2.0", fake.createInput["model"])
}

func TestDoRequestMapsOfficialSeedanceModelToChinaMobileDefault(t *testing.T) {
fake := &fakeSDK{createID: "upstream-task-1"}
adaptor := &TaskAdaptor{
newSDK: func(baseURL, apiKey, model string) (seedanceSDK, error) {
require.Equal(t, "doubao-seedance-2.0", model)
return fake, nil
},
}
payload := strings.NewReader(`{"model":"doubao-seedance-2-0-260128","content":[{"type":"text","text":"hello"}]}`)

resp, err := adaptor.DoRequest(&gin.Context{}, &relaycommon.RelayInfo{
ChannelMeta: &relaycommon.ChannelMeta{ChannelType: constant.ChannelTypeChinaMobileSeedance},
OriginModelName: "doubao-seedance-2-0-260128",
}, payload)

require.NoError(t, err)
require.Equal(t, http.StatusOK, resp.StatusCode)
require.Equal(t, "doubao-seedance-2.0", fake.createInput["model"])
}

func TestDoRequestMapsPermissionErrorToForbiddenResponse(t *testing.T) {
adaptor := &TaskAdaptor{
newSDK: func(baseURL, apiKey, model string) (seedanceSDK, error) {
return &fakeSDK{err: errors.New(`Failed to create video generation task: {"ErrorCode":"PERMISSION_ERROR","ErrorMessage":"Endpoint is not authorized"}`)}, nil
},
}
payload := strings.NewReader(`{"model":"doubao-seedance-2.0","content":[{"type":"text","text":"hello"}]}`)

resp, err := adaptor.DoRequest(&gin.Context{}, &relaycommon.RelayInfo{
ChannelMeta: &relaycommon.ChannelMeta{ChannelType: constant.ChannelTypeChinaMobileSeedance},
OriginModelName: "doubao-seedance-2.0",
}, payload)

require.NoError(t, err)
require.Equal(t, http.StatusForbidden, resp.StatusCode)
data, err := io.ReadAll(resp.Body)
require.NoError(t, err)
require.Contains(t, string(data), "PERMISSION_ERROR")
require.Contains(t, string(data), "Endpoint is not authorized")
}

func TestDoRequestMapsSensitiveContentErrorToBadRequestResponse(t *testing.T) {
adaptor := &TaskAdaptor{
newSDK: func(baseURL, apiKey, model string) (seedanceSDK, error) {
return &fakeSDK{err: errors.New(`Failed to create video generation task: {"ErrorCode":"InputVideoSensitiveContentDetected.PrivacyInformation","ErrorMessage":"input video may contain real person"}`)}, nil
},
}
payload := strings.NewReader(`{"model":"doubao-seedance-2.0","content":[{"type":"text","text":"hello"}]}`)

resp, err := adaptor.DoRequest(&gin.Context{}, &relaycommon.RelayInfo{
ChannelMeta: &relaycommon.ChannelMeta{ChannelType: constant.ChannelTypeChinaMobileSeedance},
OriginModelName: "doubao-seedance-2.0",
}, payload)

require.NoError(t, err)
require.Equal(t, http.StatusBadRequest, resp.StatusCode)
data, err := io.ReadAll(resp.Body)
require.NoError(t, err)
require.Contains(t, string(data), "InputVideoSensitiveContentDetected.PrivacyInformation")
}

func TestFetchTaskCallsSDKAndReturnsQueryMap(t *testing.T) {
fake := &fakeSDK{queryResult: map[string]interface{}{
"id": "upstream-task-1",
"status": "succeeded",
"content": map[string]interface{}{
"video_url": "https://example.test/video.mp4",
},
}}
adaptor := &TaskAdaptor{
newSDK: func(baseURL, apiKey, model string) (seedanceSDK, error) {
require.Equal(t, "https://cm.example.com/api/v3", baseURL)
require.Equal(t, "sk-test", apiKey)
require.Equal(t, "doubao-seedance-2.0", model)
return fake, nil
},
}
adaptor.Init(&relaycommon.RelayInfo{ChannelMeta: &relaycommon.ChannelMeta{ChannelType: constant.ChannelTypeChinaMobileSeedance}})

resp, err := adaptor.FetchTask("https://cm.example.com/api/v3", "sk-test", map[string]any{
"task_id": "upstream-task-1",
"model": "doubao-seedance-2.0",
}, "")

require.NoError(t, err)
require.Equal(t, "upstream-task-1", fake.queryTaskID)
data, err := io.ReadAll(resp.Body)
require.NoError(t, err)
require.Contains(t, string(data), `"video_url":"https://example.test/video.mp4"`)
}

func TestFetchTaskWithoutInitUsesDefaultTimeout(t *testing.T) {
fake := &fakeSDK{queryResult: map[string]interface{}{
"id": "upstream-task-1",
"status": "queued",
}}
adaptor := &TaskAdaptor{
newSDK: func(baseURL, apiKey, model string) (seedanceSDK, error) {
return fake, nil
},
}

resp, err := adaptor.FetchTask("https://cm.example.com/api/v3", "sk-test", map[string]any{
"task_id": "upstream-task-1",
"model": "doubao-seedance-2.0",
}, "")

require.NoError(t, err)
require.Equal(t, http.StatusOK, resp.StatusCode)
}

func TestDoResponseKeepsUpstreamIDWhenPublicTaskIDMissing(t *testing.T) {
adaptor := &TaskAdaptor{}
w := httptest.NewRecorder()
c, _ := gin.CreateTestContext(w)
resp := &http.Response{
StatusCode: http.StatusOK,
Body: io.NopCloser(strings.NewReader(`{"id":"upstream-task-1"}`)),
}

upstreamID, _, taskErr := adaptor.DoResponse(c, resp, &relaycommon.RelayInfo{
OriginModelName: "doubao-seedance-2.0",
TaskRelayInfo: &relaycommon.TaskRelayInfo{},
})

require.Nil(t, taskErr)
require.Equal(t, "upstream-task-1", upstreamID)
require.Contains(t, w.Body.String(), `"id":"upstream-task-1"`)
}

func TestParseTaskResultSuccessWithVideoURLAndUsage(t *testing.T) {
taskInfo, err := (&TaskAdaptor{}).ParseTaskResult([]byte(`{
"id":"cm-task",
"status":"succeeded",
"content":{"video_url":"https://example.com/video.mp4"},
"usage":{"completion_tokens":12,"total_tokens":34}
}`))
require.NoError(t, err)
require.Equal(t, model.TaskStatusSuccess, taskInfo.Status)
require.Equal(t, "https://example.com/video.mp4", taskInfo.Url)
require.Equal(t, 12, taskInfo.CompletionTokens)
require.Equal(t, 34, taskInfo.TotalTokens)
}

func TestParseTaskResultSuccessWithoutUsage(t *testing.T) {
taskInfo, err := (&TaskAdaptor{}).ParseTaskResult([]byte(`{
"id":"cm-task",
"status":"success",
"content":{"video_url":"https://example.com/video.mp4"}
}`))
require.NoError(t, err)
require.Equal(t, model.TaskStatusSuccess, taskInfo.Status)
require.Equal(t, 0, taskInfo.CompletionTokens)
require.Equal(t, 0, taskInfo.TotalTokens)
}

func TestParseTaskResultFailureFromErrorMap(t *testing.T) {
taskInfo, err := (&TaskAdaptor{}).ParseTaskResult([]byte(`{
"error":{"code":"bad_request","message":"invalid content"}
}`))
require.NoError(t, err)
require.Equal(t, model.TaskStatusFailure, taskInfo.Status)
require.Equal(t, "invalid content", taskInfo.Reason)
}

func TestParseTaskResultFailureFromErrorString(t *testing.T) {
taskInfo, err := (&TaskAdaptor{}).ParseTaskResult([]byte(`{
"status":"failed",
"error":"quota exceeded"
}`))
require.NoError(t, err)
require.Equal(t, model.TaskStatusFailure, taskInfo.Status)
require.Equal(t, "quota exceeded", taskInfo.Reason)
}

func newTaskRequestContext(t *testing.T, body string) *gin.Context {
t.Helper()
w := httptest.NewRecorder()
c, _ := gin.CreateTestContext(w)
c.Request = httptest.NewRequest(http.MethodPost, "/api/v3/contents/generations/tasks", strings.NewReader(body))
c.Request.Header.Set("Content-Type", "application/json")
var req relaycommon.TaskSubmitReq
require.NoError(t, common.Unmarshal([]byte(body), &req))
relaycommon.StoreTaskRequest(c, &relaycommon.RelayInfo{}, constant.TaskActionGenerate, req)
return c
}

+ 9
- 0
relay/channel/task/chinamobile_seedance/constants.go Просмотреть файл

@@ -0,0 +1,9 @@
package chinamobile_seedance

var ModelList = []string{
"doubao-seedance-2.0",
"doubao-seedance-2-0-260128",
"doubao-seedance-2-0-fast-260128",
}

var ChannelName = "ChinaMobileSeedance"

+ 81
- 0
relay/channel/task/chinamobile_seedance/sdk_client.go Просмотреть файл

@@ -0,0 +1,81 @@
package chinamobile_seedance

import (
"fmt"
"strings"
"sync"
"time"

maas "maas_seedance_sdk_1.0.0_go"
)

const (
defaultBaseURL = "https://zhenze-huhehaote.cmecloud.cn/api/v3"
defaultModel = "doubao-seedance-2.0"
defaultSubmitTimeout = 180 * time.Second
defaultQueryTimeout = 60 * time.Second
sdkAcquireTimeout = 2 * time.Second
)

var sdkSemaphore = make(chan struct{}, 4)

type seedanceSDK interface {
CreateVideoGenerationTask(data map[string]interface{}) (string, error)
QueryVideoGenerationTask(taskID string) (map[string]interface{}, error)
}

type seedanceSDKFactory func(baseURL, apiKey, model string) (seedanceSDK, error)

func newRealSDK(baseURL, apiKey, model string) (seedanceSDK, error) {
baseURL = strings.TrimRight(strings.TrimSpace(baseURL), "/")
if baseURL == "" {
baseURL = defaultBaseURL
}
model = strings.TrimSpace(model)
if model == "" {
model = defaultModel
}
return runSDKCall[seedanceSDK](defaultQueryTimeout, func() (seedanceSDK, error) {
return maas.NewMaasSeedanceClient(baseURL, apiKey, model, false)
})
}

func runSDKCall[T any](timeout time.Duration, fn func() (T, error)) (T, error) {
var zero T
select {
case sdkSemaphore <- struct{}{}:
case <-time.After(sdkAcquireTimeout):
return zero, fmt.Errorf("upstream_busy: sdk concurrency limit reached")
}

type result struct {
value T
err error
}
var releaseOnce sync.Once
release := func() {
releaseOnce.Do(func() {
<-sdkSemaphore
})
}
done := make(chan result, 1)
go func() {
res := result{}
defer func() {
if r := recover(); r != nil {
res.err = fmt.Errorf("upstream_sdk_panic: %v", r)
}
release()
done <- res
}()
res.value, res.err = fn()
}()

select {
case res := <-done:
return res.value, res.err
case <-time.After(timeout):
release()
return zero, fmt.Errorf("upstream_timeout: sdk call exceeded %s", timeout)
}
}

+ 24
- 0
relay/channel/task/doubao_tianyiyun/adaptor.go Просмотреть файл

@@ -5,6 +5,7 @@ import (
"fmt"
"io"
"net/http"
"net/url"
"strconv"
"strings"
"time"
@@ -231,6 +232,13 @@ func (a *TaskAdaptor) convertToRequestPayload(req *relaycommon.TaskSubmitReq) (*
if err := taskcommon.UnmarshalMetadata(req.Metadata, &r); err != nil {
return nil, errors.Wrap(err, "unmarshal metadata failed")
}
for _, content := range r.Content {
if !tianyiYunMediaURLIsSupported(content.ImageURL) ||
!tianyiYunMediaURLIsSupported(content.VideoURL) ||
!tianyiYunMediaURLIsSupported(content.AudioURL) {
return nil, fmt.Errorf("天翼云素材仅支持公网 URL 或 asset:// 标识")
}
}
if sec, _ := strconv.Atoi(req.Seconds); sec > 0 && r.Duration == nil {
r.Duration = lo.ToPtr(dto.IntValue(sec))
}
@@ -243,6 +251,22 @@ func (a *TaskAdaptor) convertToRequestPayload(req *relaycommon.TaskSubmitReq) (*
return &r, nil
}

func tianyiYunMediaURLIsSupported(media *MediaURL) bool {
if media == nil {
return true
}
parsedURL, err := url.Parse(strings.TrimSpace(media.URL))
if err != nil || parsedURL.Host == "" {
return false
}
switch strings.ToLower(parsedURL.Scheme) {
case "http", "https", "asset":
return true
default:
return false
}
}

func (a *TaskAdaptor) ParseTaskResult(respBody []byte) (*relaycommon.TaskInfo, error) {
resTask := responseTask{}
if err := common.Unmarshal(respBody, &resTask); err != nil {


+ 51
- 0
relay/channel/task/doubao_tianyiyun/adaptor_test.go Просмотреть файл

@@ -49,9 +49,60 @@ func TestGetModelListIncludesTianyiYunSeedanceModels(t *testing.T) {

require.Contains(t, models, "cdance2.0-0611")
require.Contains(t, models, "cdance2.0-fast-0611")
require.Contains(t, models, "cdance2.0-mini-0611")
require.Equal(t, "DoubaoVideoCompatibleTianyiYun", (&TaskAdaptor{}).GetChannelName())
}

func TestBuildRequestBodyRejectsTianyiYunDataURI(t *testing.T) {
for _, mediaType := range []string{"image_url", "video_url", "audio_url"} {
t.Run(mediaType, func(t *testing.T) {
adaptor := &TaskAdaptor{}
c := newTianyiYunTaskRequestContext(t, `{
"model":"cdance2.0-0611",
"metadata":{"content":[{"type":"`+mediaType+`","`+mediaType+`":{"url":"data:image/png;base64,AAAA"}}]}
}`)
_, err := adaptor.BuildRequestBody(c, &relaycommon.RelayInfo{ChannelMeta: &relaycommon.ChannelMeta{UpstreamModelName: "cdance2.0-0611"}})
require.ErrorContains(t, err, "asset://")
})
}
}

func TestBuildRequestBodyRejectsTianyiYunUnsupportedMediaURL(t *testing.T) {
for _, mediaURL := range []string{
"file:///tmp/input.png",
"ftp://example.test/input.png",
"aGVsbG8=",
} {
t.Run(mediaURL, func(t *testing.T) {
adaptor := &TaskAdaptor{}
c := newTianyiYunTaskRequestContext(t, `{
"model":"cdance2.0-0611",
"metadata":{"content":[{"type":"image_url","image_url":{"url":"`+mediaURL+`"}}]}
}`)
_, err := adaptor.BuildRequestBody(c, &relaycommon.RelayInfo{ChannelMeta: &relaycommon.ChannelMeta{UpstreamModelName: "cdance2.0-0611"}})
require.ErrorContains(t, err, "asset://")
})
}
}

func TestBuildRequestBodyPreservesTianyiYunAssetReference(t *testing.T) {
adaptor := &TaskAdaptor{}
c := newTianyiYunTaskRequestContext(t, `{
"model":"cdance2.0-0611",
"metadata":{"content":[
{"type":"image_url","image_url":{"url":"asset://asset-1"},"role":"reference_image"},
{"type":"video_url","video_url":{"url":"https://example.test/input.mp4"},"role":"reference_video"}
]}
}`)
body, err := adaptor.BuildRequestBody(c, &relaycommon.RelayInfo{ChannelMeta: &relaycommon.ChannelMeta{UpstreamModelName: "cdance2.0-0611"}})
require.NoError(t, err)
data, err := io.ReadAll(body)
require.NoError(t, err)
require.Contains(t, string(data), `"url":"asset://asset-1"`)
require.Contains(t, string(data), `"url":"https://example.test/input.mp4"`)
require.Less(t, strings.Index(string(data), "asset://asset-1"), strings.Index(string(data), "https://example.test/input.mp4"))
}

func TestBuildRequestBodyPreservesTianyiYunNativeFields(t *testing.T) {
adaptor := &TaskAdaptor{}
c := newTianyiYunTaskRequestContext(t, `{


+ 1
- 0
relay/channel/task/doubao_tianyiyun/constants.go Просмотреть файл

@@ -3,6 +3,7 @@ package doubao_tianyiyun
var ModelList = []string{
"cdance2.0-0611",
"cdance2.0-fast-0611",
"cdance2.0-mini-0611",
}

var ChannelName = "DoubaoVideoCompatibleTianyiYun"

+ 9
- 0
relay/helper/matrix_usage_capability.go Просмотреть файл

@@ -23,6 +23,15 @@ var matrixUsageCapabilities = []MatrixUsageCapability{
{BillingModelName: "doubao-seedance-2-0-fast-260128", UpstreamModelName: "doubao-seedance-2-0-fast-260128", ChannelType: constant.ChannelTypeDoubaoVideoCompatibleAiping},
{BillingModelName: "cdance2.0-0611", UpstreamModelName: "cdance2.0-0611", ChannelType: constant.ChannelTypeDoubaoVideoCompatibleTianyiYun},
{BillingModelName: "cdance2.0-fast-0611", UpstreamModelName: "cdance2.0-fast-0611", ChannelType: constant.ChannelTypeDoubaoVideoCompatibleTianyiYun},
{BillingModelName: "cdance2.0-mini-0611", UpstreamModelName: "cdance2.0-mini-0611", ChannelType: constant.ChannelTypeDoubaoVideoCompatibleTianyiYun},
{BillingModelName: "Doubao-Seedance-2.0-mini", UpstreamModelName: "cdance2.0-mini-0611", ChannelType: constant.ChannelTypeDoubaoVideoCompatibleTianyiYun},
{BillingModelName: "doubao-seedance-2-0-260128", UpstreamModelName: "doubao-seedance-2-0-260128", ChannelType: constant.ChannelTypeChinaMobileSeedance},
{BillingModelName: "doubao-seedance-2-0-260128", UpstreamModelName: "doubao-seedance-2.0", ChannelType: constant.ChannelTypeChinaMobileSeedance},
{BillingModelName: "doubao-seedance-2-0-fast-260128", UpstreamModelName: "doubao-seedance-2-0-fast-260128", ChannelType: constant.ChannelTypeChinaMobileSeedance},
{BillingModelName: "doubao-seedance-2-0-fast-260128", UpstreamModelName: "doubao-seedance-2.0", ChannelType: constant.ChannelTypeChinaMobileSeedance},
{BillingModelName: "doubao-seedance-2.0", UpstreamModelName: "doubao-seedance-2.0", ChannelType: constant.ChannelTypeChinaMobileSeedance},
{BillingModelName: "doubao-seedance-2-0-mini-260615", UpstreamModelName: "doubao-seedance-2-0-mini-260615", ChannelType: constant.ChannelTypeChinaMobileSeedance},
{BillingModelName: "doubao-seedance-2-0-mini-260615", UpstreamModelName: "doubao-seedance-2.0", ChannelType: constant.ChannelTypeChinaMobileSeedance},
// Keep in sync with SupportsAnyMatrixUsageBillingModel in setting/ratio_setting/model_pricing.go
{BillingModelName: "kling-v1", UpstreamModelName: "kling-v1", ChannelType: constant.ChannelTypeKlingAiping},
{BillingModelName: "kling-v1", UpstreamModelName: "Kling-V1", ChannelType: constant.ChannelTypeKlingAiping},


+ 18
- 0
relay/helper/matrix_usage_capability_test.go Просмотреть файл

@@ -29,6 +29,24 @@ func TestSupportsMatrixUsageBilling_TianyiYunSeedanceChannel(t *testing.T) {
require.False(t, SupportsMatrixUsageBilling("cdance2.0-0611", constant.ChannelTypeDoubaoVideoCompatibleAiping, "cdance2.0-0611"))
}

func TestSupportsMatrixUsageBilling_TianyiYunMini(t *testing.T) {
require.True(t, SupportsMatrixUsageBilling("cdance2.0-mini-0611", constant.ChannelTypeDoubaoVideoCompatibleTianyiYun, "cdance2.0-mini-0611"))
require.False(t, SupportsMatrixUsageBilling("cdance2.0-mini-0611", constant.ChannelTypeDoubaoVideoCompatibleAiping, "cdance2.0-mini-0611"))
require.False(t, SupportsMatrixUsageBilling("cdance2.0-mini-0611", constant.ChannelTypeDoubaoVideoCompatibleTianyiYun, "other-model"))
require.True(t, SupportsMatrixUsageBilling("Doubao-Seedance-2.0-mini", constant.ChannelTypeDoubaoVideoCompatibleTianyiYun, "cdance2.0-mini-0611"))
require.False(t, SupportsMatrixUsageBilling("Doubao-Seedance-2.0-mini", constant.ChannelTypeDoubaoVideoCompatibleAiping, "cdance2.0-mini-0611"))
require.False(t, SupportsMatrixUsageBilling("Doubao-Seedance-2.0-mini", constant.ChannelTypeDoubaoVideoCompatibleTianyiYun, "other-model"))
}

func TestSupportsMatrixUsageBilling_ChinaMobileSeedanceChannel(t *testing.T) {
require.True(t, SupportsMatrixUsageBilling("doubao-seedance-2-0-260128", constant.ChannelTypeChinaMobileSeedance, "doubao-seedance-2-0-260128"))
require.True(t, SupportsMatrixUsageBilling("doubao-seedance-2-0-260128", constant.ChannelTypeChinaMobileSeedance, "doubao-seedance-2.0"))
require.True(t, SupportsMatrixUsageBilling("doubao-seedance-2-0-fast-260128", constant.ChannelTypeChinaMobileSeedance, "doubao-seedance-2-0-fast-260128"))
require.True(t, SupportsMatrixUsageBilling("doubao-seedance-2-0-fast-260128", constant.ChannelTypeChinaMobileSeedance, "doubao-seedance-2.0"))
require.True(t, SupportsMatrixUsageBilling("doubao-seedance-2.0", constant.ChannelTypeChinaMobileSeedance, "doubao-seedance-2.0"))
require.False(t, SupportsMatrixUsageBilling("doubao-seedance-2-0-260128", constant.ChannelTypeDoubaoVideoCompatibleTianyiYun, "doubao-seedance-2.0"))
}

func TestSupportsMatrixUsageBilling_KlingAipingVideoModels(t *testing.T) {
cases := []struct {
billingModel string


+ 3
- 0
relay/relay_adaptor.go Просмотреть файл

@@ -31,6 +31,7 @@ import (
"github.com/QuantumNous/new-api/relay/channel/siliconflow"
"github.com/QuantumNous/new-api/relay/channel/submodel"
taskali "github.com/QuantumNous/new-api/relay/channel/task/ali"
taskchinamobileseedance "github.com/QuantumNous/new-api/relay/channel/task/chinamobile_seedance"
taskdoubao "github.com/QuantumNous/new-api/relay/channel/task/doubao"
taskdoubaoaiping "github.com/QuantumNous/new-api/relay/channel/task/doubao_aiping"
taskdoubaotianyiyun "github.com/QuantumNous/new-api/relay/channel/task/doubao_tianyiyun"
@@ -160,6 +161,8 @@ func GetTaskAdaptor(platform constant.TaskPlatform) channel.TaskAdaptor {
return &taskdoubaoaiping.TaskAdaptor{}
case constant.ChannelTypeDoubaoVideoCompatibleTianyiYun:
return &taskdoubaotianyiyun.TaskAdaptor{}
case constant.ChannelTypeChinaMobileSeedance:
return &taskchinamobileseedance.TaskAdaptor{}
case constant.ChannelTypeKlingAiping:
return &klingaiping.TaskAdaptor{}
case constant.ChannelTypeSora, constant.ChannelTypeOpenAI:


+ 15
- 0
relay/relay_adaptor_chinamobile_test.go Просмотреть файл

@@ -0,0 +1,15 @@
package relay

import (
"strconv"
"testing"

"github.com/QuantumNous/new-api/constant"
"github.com/stretchr/testify/require"
)

func TestGetTaskAdaptorChinaMobileSeedance(t *testing.T) {
adaptor := GetTaskAdaptor(constant.TaskPlatform(strconv.Itoa(constant.ChannelTypeChinaMobileSeedance)))
require.NotNil(t, adaptor)
require.Equal(t, "ChinaMobileSeedance", adaptor.GetChannelName())
}

+ 5
- 0
relay/relay_task.go Просмотреть файл

@@ -406,9 +406,14 @@ func tryRealtimeFetch(task *model.Task, isOpenAIVideoAPI bool) []byte {
return nil
}

fetchModel := task.Properties.UpstreamModelName
if strings.TrimSpace(fetchModel) == "" {
fetchModel = task.Properties.OriginModelName
}
resp, err := adaptor.FetchTask(baseURL, channelModel.Key, map[string]any{
"task_id": task.GetUpstreamTaskID(),
"action": task.Action,
"model": fetchModel,
}, proxy)
if err != nil || resp == nil {
return nil


+ 2
- 0
router/api-router.go Просмотреть файл

@@ -136,6 +136,8 @@ func SetApiRouter(router *gin.Engine) {
adminRoute.GET("/:id/oauth/bindings", controller.GetUserOAuthBindingsByAdmin)
adminRoute.DELETE("/:id/oauth/bindings/:provider_id", controller.UnbindCustomOAuthByAdmin)
adminRoute.DELETE("/:id/bindings/:binding_type", controller.AdminClearUserBinding)
adminRoute.GET("/:id/video-channel-bindings", controller.GetUserVideoChannelBindings)
adminRoute.PUT("/:id/video-channel-bindings", controller.SetUserVideoChannelBindings)
adminRoute.GET("/:id", controller.GetUser)
adminRoute.POST("/", controller.CreateUser)
adminRoute.POST("/manage", controller.ManageUser)


+ 101
- 0
service/asset_adapter.go Просмотреть файл

@@ -0,0 +1,101 @@
package service

import (
"context"
"net/http"
"sort"
"strings"

"github.com/QuantumNous/new-api/constant"
"github.com/QuantumNous/new-api/model"
)

const (
AssetErrorInvalidRequest = "invalid_request_error"
AssetErrorChannelNotFound = "asset_channel_not_found"
AssetErrorBindingInvalid = "asset_channel_binding_invalid"
AssetErrorOperationNotSupported = "asset_operation_not_supported"
AssetErrorNotFound = "asset_not_found"
AssetErrorUpstream = "upstream_error"
AssetErrorServer = "server_error"
)

type AssetError struct {
Type string
Message string
HTTPStatus int
}

func newAssetError(errType string, message string, status int) *AssetError {
if status == 0 {
status = http.StatusBadGateway
}
return &AssetError{Type: errType, Message: message, HTTPStatus: status}
}

type AssetRequest struct {
Action AssetActionSpec
Version string
Body map[string]any
RawBody []byte
}

type AssetUpstreamResponse struct {
StatusCode int
Header http.Header
Body []byte
}

type AssetAdapter interface {
Name() string
Supports(operation AssetOperation) bool
DoAssetRequest(ctx context.Context, channel *model.Channel, req AssetRequest) (*AssetUpstreamResponse, *AssetError)
}

var assetAdapters = map[int]AssetAdapter{}

func init() {
registerDefaultAssetAdapters()
}

func registerDefaultAssetAdapters() {
assetAdapters[constant.ChannelTypeChinaMobileSeedance] = NewChinaMobileAssetAdapter()
assetAdapters[constant.ChannelTypeDoubaoVideoCompatibleAiping] = NewCompatibleAssetAdapter("aiping_asset", []AssetOperation{
AssetOperationAssetCreate,
AssetOperationAssetList,
AssetOperationAssetGet,
AssetOperationAssetUpdate,
AssetOperationAssetDelete,
})
assetAdapters[constant.ChannelTypeDoubaoVideoCompatibleTianyiYun] = NewTianyiYunAssetAdapter()
}

func GetAssetAdapter(channelType int) (AssetAdapter, bool) {
adapter, ok := assetAdapters[channelType]
return adapter, ok
}

func OverrideAssetAdapterForTest(channelType int, adapter AssetAdapter) func() {
oldAdapter, hadOldAdapter := assetAdapters[channelType]
assetAdapters[channelType] = adapter
return func() {
if hadOldAdapter {
assetAdapters[channelType] = oldAdapter
} else {
delete(assetAdapters, channelType)
}
}
}

func RegisteredAssetChannelTypes() []int {
types := make([]int, 0, len(assetAdapters))
for channelType := range assetAdapters {
types = append(types, channelType)
}
sort.Ints(types)
return types
}

func assetChannelHasKey(channel *model.Channel) bool {
return channel != nil && strings.TrimSpace(channel.Key) != "" && len(channel.GetKeys()) > 0
}

+ 28
- 0
service/asset_adapter_test.go Просмотреть файл

@@ -0,0 +1,28 @@
package service

import (
"testing"

"github.com/stretchr/testify/assert"
)

func resetAssetAdapterRegistryForTest(t interface{ Cleanup(func()) }) {
oldRegistry := assetAdapters
assetAdapters = map[int]AssetAdapter{}
t.Cleanup(func() {
assetAdapters = oldRegistry
})
}

func RegisterAssetAdapterForTest(channelType int, adapter AssetAdapter) {
assetAdapters[channelType] = adapter
}

func TestRegisteredAssetChannelTypesAreSorted(t *testing.T) {
resetAssetAdapterRegistryForTest(t)
RegisterAssetAdapterForTest(61, fakeAssetAdapter{})
RegisterAssetAdapterForTest(7, fakeAssetAdapter{})
RegisterAssetAdapterForTest(16, fakeAssetAdapter{})

assert.Equal(t, []int{7, 16, 61}, RegisteredAssetChannelTypes())
}

+ 758
- 0
service/asset_chinamobile.go Просмотреть файл

@@ -0,0 +1,758 @@
package service

import (
"context"
"errors"
"fmt"
"net/http"
"net/url"
"os"
"strings"

"github.com/QuantumNous/new-api/common"
"github.com/QuantumNous/new-api/model"
cmerrs "gitlab.ecloud.com/ecloud/ecloudsdkcore/errs"
cmmodel "gitlab.ecloud.com/ecloud/ecloudsdkmaas/model"
)

const defaultChinaMobileAssetBaseURL = "https://ecloud.10086.cn"
const defaultChinaMobileAssetPoolID = "CIDC-CORE-00"
const chinaMobileAssetAKEnv = "CHINAMOBILE_ASSET_AK"
const chinaMobileAssetSKEnv = "CHINAMOBILE_ASSET_SK"
const chinaMobileAssetPoolIDEnv = "CHINAMOBILE_ASSET_POOL_ID"

type ChinaMobileAssetAdapter struct {
newClient func(credential chinaMobileAssetCredential) chinaMobileAssetSDKClient
}

func NewChinaMobileAssetAdapter() AssetAdapter {
return &ChinaMobileAssetAdapter{newClient: newChinaMobileAssetSDKClient}
}

func (a *ChinaMobileAssetAdapter) Name() string {
return "chinamobile_asset"
}

func (a *ChinaMobileAssetAdapter) Supports(operation AssetOperation) bool {
switch operation {
case AssetOperationAssetCreate,
AssetOperationAssetList,
AssetOperationAssetGet,
AssetOperationAssetUpdate,
AssetOperationAssetDelete,
AssetOperationAssetGroupCreate,
AssetOperationAssetGroupList,
AssetOperationAssetGroupGet,
AssetOperationAssetGroupUpdate,
AssetOperationAssetGroupDelete:
return true
default:
return false
}
}

func (a *ChinaMobileAssetAdapter) DoAssetRequest(ctx context.Context, channel *model.Channel, req AssetRequest) (*AssetUpstreamResponse, *AssetError) {
if !a.Supports(req.Action.Operation) {
return nil, newAssetError(AssetErrorOperationNotSupported, fmt.Sprintf("asset operation %s is not supported", req.Action.Operation), http.StatusBadRequest)
}
if err := validateChinaMobileCompatibility(req.Body); err != nil {
return nil, err
}
if err := validateChinaMobileAssetRequest(req); err != nil {
return nil, err
}

credential, err := chinaMobileAssetCredentialFromChannel(channel)
if err != nil {
return nil, newAssetError(AssetErrorInvalidRequest, err.Error(), http.StatusBadRequest)
}
newClient := a.newClient
if newClient == nil {
newClient = newChinaMobileAssetSDKClient
}
result, err := callChinaMobileAssetSDK(ctx, newClient(credential), req)
if err != nil {
if assetErr := classifyChinaMobileAssetSDKError(err); assetErr != nil {
return nil, assetErr
}
return nil, newAssetError(AssetErrorUpstream, err.Error(), http.StatusBadGateway)
}
normalized, assetErr := normalizeChinaMobileAssetSDKResponse(req.Action, req.Version, result)
if assetErr != nil {
return nil, assetErr
}
return &AssetUpstreamResponse{StatusCode: http.StatusOK, Header: http.Header{}, Body: normalized}, nil
}

func joinChinaMobileAssetURL(baseURL string, path string) (string, error) {
u, err := url.Parse(baseURL)
if err != nil {
return "", err
}
u.Path = strings.TrimRight(u.Path, "/") + path
u.RawQuery = ""
u.Fragment = ""
return u.String(), nil
}

type chinaMobileAssetCredential struct {
AK string `json:"ak"`
SK string `json:"sk"`
PoolID string `json:"pool_id"`
}

func chinaMobileAssetCredentialFromChannel(channel *model.Channel) (chinaMobileAssetCredential, error) {
if channel == nil {
return chinaMobileAssetCredential{}, fmt.Errorf("China Mobile asset channel is required")
}
credential, err := model.GetChannelAssetCredential(channel.Id)
if err != nil {
return chinaMobileAssetCredential{}, err
}
if credential == nil {
legacyCredential, legacyErr := normalizeChinaMobileAssetCredential(chinaMobileAssetCredential{
AK: os.Getenv(chinaMobileAssetAKEnv),
SK: os.Getenv(chinaMobileAssetSKEnv),
PoolID: os.Getenv(chinaMobileAssetPoolIDEnv),
})
if legacyErr != nil {
return chinaMobileAssetCredential{}, fmt.Errorf("该移动云渠道未配置素材凭证")
}
common.SysLog(fmt.Sprintf("using legacy China Mobile asset credentials for channel %d; configure channel asset credentials before removing environment fallback", channel.Id))
return legacyCredential, nil
}
return normalizeChinaMobileAssetCredential(chinaMobileAssetCredential{
AK: credential.AccessKey,
SK: credential.SecretKey,
PoolID: credential.PoolID,
})
}

func normalizeChinaMobileAssetCredential(credential chinaMobileAssetCredential) (chinaMobileAssetCredential, error) {
credential.AK = strings.TrimSpace(credential.AK)
credential.SK = strings.TrimSpace(credential.SK)
credential.PoolID = strings.TrimSpace(credential.PoolID)
if credential.AK == "" || credential.SK == "" {
return chinaMobileAssetCredential{}, fmt.Errorf("China Mobile asset AccessKey and SecretKey are required")
}
if credential.PoolID == "" {
credential.PoolID = defaultChinaMobileAssetPoolID
}
return credential, nil
}

func callChinaMobileAssetSDK(ctx context.Context, client chinaMobileAssetSDKClient, req AssetRequest) (any, error) {
type sdkResult struct {
value any
err error
}
resultCh := make(chan sdkResult, 1)
go func() {
value, err := executeChinaMobileAssetSDK(client, req)
resultCh <- sdkResult{value: value, err: err}
}()
select {
case <-ctx.Done():
return nil, ctx.Err()
case result := <-resultCh:
return result.value, result.err
}
}

func executeChinaMobileAssetSDK(client chinaMobileAssetSDKClient, req AssetRequest) (any, error) {
switch req.Action.Operation {
case AssetOperationAssetCreate:
return client.CreateAsset(&cmmodel.CreateAssetRequest{CreateAssetBody: newChinaMobileCreateAssetBody(req.Body)})
case AssetOperationAssetList:
return client.ListAssets(&cmmodel.ListAssetsRequest{ListAssetsBody: newChinaMobileListAssetsBody(req.Body)})
case AssetOperationAssetGet:
id, err := requiredAssetString(req.Body, "Id")
if err != nil {
return nil, err
}
return client.GetAsset(&cmmodel.GetAssetRequest{GetAssetPath: (&cmmodel.GetAssetPath{}).SetAssetId(id)})
case AssetOperationAssetUpdate:
id, err := requiredAssetString(req.Body, "Id")
if err != nil {
return nil, err
}
return client.UpdateAsset(&cmmodel.UpdateAssetRequest{
UpdateAssetPath: (&cmmodel.UpdateAssetPath{}).SetAssetId(id),
UpdateAssetBody: newChinaMobileUpdateAssetBody(req.Body),
})
case AssetOperationAssetDelete:
id, err := requiredAssetString(req.Body, "Id")
if err != nil {
return nil, err
}
return client.DeleteAsset(&cmmodel.DeleteAssetRequest{DeleteAssetPath: (&cmmodel.DeleteAssetPath{}).SetAssetId(id)})
case AssetOperationAssetGroupCreate:
return client.CreateAssetGroup(&cmmodel.CreateAssetGroupRequest{CreateAssetGroupBody: newChinaMobileCreateAssetGroupBody(req.Body)})
case AssetOperationAssetGroupList:
return client.ListAssetGroups(&cmmodel.ListAssetGroupsRequest{ListAssetGroupsBody: newChinaMobileListAssetGroupsBody(req.Body)})
case AssetOperationAssetGroupGet:
id, err := requiredAssetString(req.Body, "Id")
if err != nil {
return nil, err
}
return client.GetAssetGroup(&cmmodel.GetAssetGroupRequest{GetAssetGroupPath: (&cmmodel.GetAssetGroupPath{}).SetGroupId(id)})
case AssetOperationAssetGroupUpdate:
id, err := requiredAssetString(req.Body, "Id")
if err != nil {
return nil, err
}
return client.UpdateAssetGroup(&cmmodel.UpdateAssetGroupRequest{
UpdateAssetGroupPath: (&cmmodel.UpdateAssetGroupPath{}).SetGroupId(id),
UpdateAssetGroupBody: newChinaMobileUpdateAssetGroupBody(req.Body),
})
case AssetOperationAssetGroupDelete:
id, err := requiredAssetString(req.Body, "Id")
if err != nil {
return nil, err
}
return client.DeleteAssetGroup(&cmmodel.DeleteAssetGroupRequest{DeleteAssetGroupPath: (&cmmodel.DeleteAssetGroupPath{}).SetGroupId(id)})
default:
return nil, fmt.Errorf("unsupported asset operation %s", req.Action.Operation)
}
}

func isChinaMobileAssetLocalValidationError(err error) bool {
return err != nil && strings.Contains(err.Error(), " is required")
}

func classifyChinaMobileAssetSDKError(err error) *AssetError {
if err == nil {
return nil
}
if errors.Is(err, context.Canceled) || errors.Is(err, context.DeadlineExceeded) {
return newAssetError(AssetErrorUpstream, err.Error(), http.StatusGatewayTimeout)
}
if isChinaMobileAssetLocalValidationError(err) {
return newAssetError(AssetErrorInvalidRequest, err.Error(), http.StatusBadRequest)
}
var responseErr *cmerrs.ServerResponseError
if errors.As(err, &responseErr) && responseErr.Code == http.StatusBadRequest {
message := strings.TrimSpace(responseErr.Body)
if message == "" {
message = responseErr.Error()
}
return newAssetError(AssetErrorInvalidRequest, message, http.StatusBadRequest)
}
return nil
}

func validateChinaMobileAssetRequest(req AssetRequest) *AssetError {
invalid := func(message string) *AssetError {
return newAssetError(AssetErrorInvalidRequest, message, http.StatusBadRequest)
}
requireString := func(key string) *AssetError {
if strings.TrimSpace(stringValue(req.Body, key)) == "" {
return invalid(key + " is required")
}
return nil
}
validatePage := func() *AssetError {
pageNumber := intValue(req.Body, "PageNumber", 1)
pageSize := intValue(req.Body, "PageSize", 10)
if pageNumber < 1 {
return invalid("PageNumber must be greater than or equal to 1")
}
if pageSize < 1 || pageSize > 999999 {
return invalid("PageSize must be between 1 and 999999")
}
return nil
}

switch req.Action.Operation {
case AssetOperationAssetCreate:
for _, key := range []string{"GroupId", "Name", "URL", "AssetType"} {
if err := requireString(key); err != nil {
return err
}
}
if len([]rune(stringValue(req.Body, "Name"))) > 64 {
return invalid("Name must not exceed 64 characters")
}
assetURL, err := url.ParseRequestURI(strings.TrimSpace(stringValue(req.Body, "URL")))
if err != nil || (assetURL.Scheme != "http" && assetURL.Scheme != "https") || assetURL.Host == "" {
return invalid("URL must be a valid public HTTP or HTTPS URL")
}
switch stringValue(req.Body, "AssetType") {
case "Image", "Video", "Audio":
default:
return invalid("AssetType must be one of Image, Video, Audio")
}
case AssetOperationAssetList:
if err := validatePage(); err != nil {
return err
}
groupType := strings.TrimSpace(stringValue(mapValue(req.Body, "Filter"), "GroupType"))
if groupType == "" {
return invalid("Filter.GroupType is required")
}
if groupType != "AIGC" && groupType != "LivenessFace" {
return invalid("Filter.GroupType must be AIGC or LivenessFace")
}
case AssetOperationAssetGroupCreate:
if stringValue(req.Body, "GroupType") != "AIGC" {
return invalid("GroupType must be AIGC")
}
if len([]rune(stringValue(req.Body, "Name"))) > 64 {
return invalid("Name must not exceed 64 characters")
}
if len([]rune(stringValue(req.Body, "Description"))) > 300 {
return invalid("Description must not exceed 300 characters")
}
case AssetOperationAssetGroupList:
return validatePage()
case AssetOperationAssetUpdate:
if err := requireString("Id"); err != nil {
return err
}
if len([]rune(stringValue(req.Body, "Name"))) > 64 {
return invalid("Name must not exceed 64 characters")
}
case AssetOperationAssetGroupUpdate:
if err := requireString("Id"); err != nil {
return err
}
if len([]rune(stringValue(req.Body, "Name"))) > 64 {
return invalid("Name must not exceed 64 characters")
}
if len([]rune(stringValue(req.Body, "Description"))) > 300 {
return invalid("Description must not exceed 300 characters")
}
case AssetOperationAssetGet, AssetOperationAssetDelete, AssetOperationAssetGroupGet, AssetOperationAssetGroupDelete:
return requireString("Id")
}
return nil
}

func newChinaMobileCreateAssetBody(body map[string]any) *cmmodel.CreateAssetBody {
assetType := cmmodel.CreateAssetBodyAssetTypeEnum(stringValue(body, "AssetType"))
result := &cmmodel.CreateAssetBody{}
result.SetGroupId(stringValue(body, "GroupId"))
result.SetAssetName(stringValue(body, "Name"))
result.SetAssetUrl(stringValue(body, "URL"))
if strings.TrimSpace(string(assetType)) != "" {
result.SetAssetType(assetType)
}
return result
}

func newChinaMobileListAssetsBody(body map[string]any) *cmmodel.ListAssetsBody {
mapped := mapChinaMobileListAssets(body)
result := &cmmodel.ListAssetsBody{}
result.SetPageNo(int32(intValue(mapped, "pageNo", 1)))
result.SetPageSize(int32(intValue(mapped, "pageSize", 10)))
if value := stringValue(mapped, "groupType"); value != "" {
result.SetGroupType(value)
}
if value := stringValue(mapped, "assetName"); value != "" {
result.SetAssetName(value)
}
if values := stringArrayValue(mapped, "groupIds"); len(values) > 0 {
result.SetGroupIds(values)
}
if values := stringArrayValue(mapped, "statuses"); len(values) > 0 {
result.SetStatuses(values)
}
return result
}

func newChinaMobileUpdateAssetBody(body map[string]any) *cmmodel.UpdateAssetBody {
result := &cmmodel.UpdateAssetBody{}
if value := stringValue(body, "Name"); value != "" {
result.SetAssetName(value)
}
return result
}

func newChinaMobileCreateAssetGroupBody(body map[string]any) *cmmodel.CreateAssetGroupBody {
result := &cmmodel.CreateAssetGroupBody{}
if value := stringValue(body, "GroupType"); value != "" {
result.SetGroupType(value)
}
if value := stringValue(body, "Name"); value != "" {
result.SetGroupName(value)
}
if value := stringValue(body, "Description"); value != "" {
result.SetDescription(value)
}
return result
}

func newChinaMobileListAssetGroupsBody(body map[string]any) *cmmodel.ListAssetGroupsBody {
mapped := mapChinaMobileListAssetGroups(body)
result := &cmmodel.ListAssetGroupsBody{}
result.SetPageNo(int32(intValue(mapped, "pageNo", 1)))
result.SetPageSize(int32(intValue(mapped, "pageSize", 10)))
if value := stringValue(mapped, "groupType"); value != "" {
result.SetGroupType(value)
}
if value := stringValue(mapped, "groupName"); value != "" {
result.SetGroupName(value)
}
if values := stringArrayValue(mapped, "groupIds"); len(values) > 0 {
result.SetGroupIds(values)
}
return result
}

func newChinaMobileUpdateAssetGroupBody(body map[string]any) *cmmodel.UpdateAssetGroupBody {
result := &cmmodel.UpdateAssetGroupBody{}
if value := stringValue(body, "Name"); value != "" {
result.SetGroupName(value)
}
if value := stringValue(body, "Description"); value != "" {
result.SetDescription(value)
}
return result
}

func buildChinaMobileAssetRequest(req AssetRequest) (string, string, map[string]any, error) {
switch req.Action.Operation {
case AssetOperationAssetCreate:
return http.MethodPost, "/api/openapi-maas/exp/aicc/v2/asset", mapChinaMobileCreateAsset(req.Body), nil
case AssetOperationAssetList:
return http.MethodPost, "/api/openapi-maas/exp/aicc/v2/asset/query", mapChinaMobileListAssets(req.Body), nil
case AssetOperationAssetGet:
id, err := requiredAssetString(req.Body, "Id")
return http.MethodGet, "/api/openapi-maas/exp/aicc/v2/asset/" + url.PathEscape(id), nil, err
case AssetOperationAssetUpdate:
id, err := requiredAssetString(req.Body, "Id")
return http.MethodPut, "/api/openapi-maas/exp/aicc/v2/asset/" + url.PathEscape(id), mapChinaMobileUpdateAsset(req.Body), err
case AssetOperationAssetDelete:
id, err := requiredAssetString(req.Body, "Id")
return http.MethodDelete, "/api/openapi-maas/exp/aicc/v2/asset/" + url.PathEscape(id), nil, err
case AssetOperationAssetGroupCreate:
return http.MethodPost, "/api/openapi-maas/exp/aicc/v2/asset-group", mapChinaMobileCreateAssetGroup(req.Body), nil
case AssetOperationAssetGroupList:
return http.MethodPost, "/api/openapi-maas/exp/aicc/v2/asset-group/query", mapChinaMobileListAssetGroups(req.Body), nil
case AssetOperationAssetGroupGet:
id, err := requiredAssetString(req.Body, "Id")
return http.MethodGet, "/api/openapi-maas/exp/aicc/v2/asset-group/" + url.PathEscape(id), nil, err
case AssetOperationAssetGroupUpdate:
id, err := requiredAssetString(req.Body, "Id")
return http.MethodPut, "/api/openapi-maas/exp/aicc/v2/asset-group/" + url.PathEscape(id), mapChinaMobileUpdateAssetGroup(req.Body), err
case AssetOperationAssetGroupDelete:
id, err := requiredAssetString(req.Body, "Id")
return http.MethodDelete, "/api/openapi-maas/exp/aicc/v2/asset-group/" + url.PathEscape(id), nil, err
default:
return "", "", nil, fmt.Errorf("unsupported asset operation %s", req.Action.Operation)
}
}

func mapChinaMobileCreateAsset(body map[string]any) map[string]any {
return compactAssetMap(map[string]any{
"groupId": stringValue(body, "GroupId"),
"assetName": stringValue(body, "Name"),
"assetUrl": stringValue(body, "URL"),
"assetType": stringValue(body, "AssetType"),
})
}

func mapChinaMobileListAssets(body map[string]any) map[string]any {
filter := mapValue(body, "Filter")
mapped := map[string]any{
"pageNo": intValue(body, "PageNumber", 1),
"pageSize": intValue(body, "PageSize", 10),
}
copyStringArray(mapped, "groupIds", filter, "GroupIds")
copyString(mapped, "groupType", filter, "GroupType")
copyString(mapped, "assetName", filter, "Name")
copyStatusArray(mapped, "statuses", filter, "Statuses")
return compactAssetMap(mapped)
}

func mapChinaMobileUpdateAsset(body map[string]any) map[string]any {
return compactAssetMap(map[string]any{"assetName": stringValue(body, "Name")})
}

func mapChinaMobileCreateAssetGroup(body map[string]any) map[string]any {
return compactAssetMap(map[string]any{
"groupType": stringValue(body, "GroupType"),
"groupName": stringValue(body, "Name"),
"description": stringValue(body, "Description"),
})
}

func mapChinaMobileListAssetGroups(body map[string]any) map[string]any {
filter := mapValue(body, "Filter")
mapped := map[string]any{
"pageNo": intValue(body, "PageNumber", 1),
"pageSize": intValue(body, "PageSize", 10),
}
copyString(mapped, "groupType", filter, "GroupType")
copyString(mapped, "groupName", filter, "Name")
copyStringArray(mapped, "groupIds", filter, "GroupIds")
return compactAssetMap(mapped)
}

func mapChinaMobileUpdateAssetGroup(body map[string]any) map[string]any {
return compactAssetMap(map[string]any{
"groupName": stringValue(body, "Name"),
"description": stringValue(body, "Description"),
})
}

func normalizeChinaMobileAssetResponse(spec AssetActionSpec, version string, data []byte) ([]byte, *AssetError) {
var payload struct {
RequestID string `json:"requestId"`
State string `json:"state"`
ErrorCode string `json:"errorCode"`
ErrorMessage string `json:"errorMessage"`
Body any `json:"body"`
}
if err := common.Unmarshal(data, &payload); err != nil {
return nil, newAssetError(AssetErrorUpstream, err.Error(), http.StatusBadGateway)
}
if strings.EqualFold(payload.State, "ERROR") {
message := strings.TrimSpace(payload.ErrorMessage)
if message == "" {
message = payload.ErrorCode
}
return nil, newAssetError(AssetErrorUpstream, message, http.StatusBadGateway)
}
if !strings.EqualFold(payload.State, "OK") {
return nil, newAssetError(AssetErrorUpstream, "China Mobile asset response has invalid state", http.StatusBadGateway)
}
if spec.Delete {
deleted, ok := payload.Body.(bool)
if !ok || !deleted {
return nil, newAssetError(AssetErrorUpstream, "China Mobile asset delete was not confirmed", http.StatusBadGateway)
}
}
var result any
if !spec.Delete {
result = normalizeChinaMobileResult(spec.Operation, payload.Body)
}
body, err := buildAssetSuccessResponseWithRequestID(spec.Action, version, payload.RequestID, result)
if err != nil {
return nil, newAssetError(AssetErrorServer, err.Error(), http.StatusInternalServerError)
}
return body, nil
}

func normalizeChinaMobileAssetSDKResponse(spec AssetActionSpec, version string, response any) ([]byte, *AssetError) {
data, err := common.Marshal(response)
if err != nil {
return nil, newAssetError(AssetErrorServer, err.Error(), http.StatusInternalServerError)
}
return normalizeChinaMobileAssetResponse(spec, version, data)
}

func normalizeChinaMobileResult(operation AssetOperation, body any) any {
raw, ok := body.(map[string]any)
if !ok {
return body
}
switch operation {
case AssetOperationAssetList:
return map[string]any{
"Items": normalizeChinaMobileItems(raw["data"], false),
"TotalCount": raw["total"],
}
case AssetOperationAssetGroupList:
return map[string]any{
"Items": normalizeChinaMobileItems(raw["data"], true),
"TotalCount": raw["total"],
}
case AssetOperationAssetGroupCreate, AssetOperationAssetGroupGet, AssetOperationAssetGroupUpdate:
return normalizeChinaMobileAssetGroup(raw)
default:
return normalizeChinaMobileAsset(raw)
}
}

func normalizeChinaMobileItems(value any, group bool) []any {
items, ok := value.([]any)
if !ok {
return nil
}
normalized := make([]any, 0, len(items))
for _, item := range items {
raw, ok := item.(map[string]any)
if !ok {
continue
}
if group {
normalized = append(normalized, normalizeChinaMobileAssetGroup(raw))
} else {
normalized = append(normalized, normalizeChinaMobileAsset(raw))
}
}
return normalized
}

func normalizeChinaMobileAsset(raw map[string]any) map[string]any {
return compactAssetMap(map[string]any{
"Id": raw["assetId"],
"GroupId": raw["groupId"],
"Name": raw["assetName"],
"AssetType": raw["assetType"],
"URL": raw["assetUrl"],
"Status": chinaMobileStatusToOfficial(fmt.Sprint(raw["status"])),
"ErrorMessage": raw["errorMessage"],
"CreatedAt": raw["createdTime"],
"UpdatedAt": raw["updatedTime"],
})
}

func normalizeChinaMobileAssetGroup(raw map[string]any) map[string]any {
return compactAssetMap(map[string]any{
"Id": raw["groupId"],
"GroupId": raw["groupId"],
"GroupType": raw["groupType"],
"Name": raw["groupName"],
"Description": raw["description"],
"CreatedAt": raw["createdTime"],
"UpdatedAt": raw["updatedTime"],
})
}

func validateChinaMobileCompatibility(body map[string]any) *AssetError {
if projectName := strings.TrimSpace(stringValue(body, "ProjectName")); projectName != "" && projectName != "default" {
return newAssetError(AssetErrorOperationNotSupported, "China Mobile asset API does not support non-default ProjectName", http.StatusBadRequest)
}
if sortBy := strings.TrimSpace(stringValue(body, "SortBy")); sortBy != "" && sortBy != "CreateTime" {
return newAssetError(AssetErrorOperationNotSupported, "China Mobile asset API does not support custom SortBy", http.StatusBadRequest)
}
if sortOrder := strings.TrimSpace(stringValue(body, "SortOrder")); sortOrder != "" && !strings.EqualFold(sortOrder, "Desc") {
return newAssetError(AssetErrorOperationNotSupported, "China Mobile asset API does not support custom SortOrder", http.StatusBadRequest)
}
return nil
}

func chinaMobileStatusToOfficial(status string) string {
switch strings.ToUpper(strings.TrimSpace(status)) {
case "PROCESSING":
return "Processing"
case "ACTIVE":
return "Active"
case "FAILED":
return "Failed"
default:
return status
}
}

func officialStatusToChinaMobile(status string) string {
switch strings.ToLower(strings.TrimSpace(status)) {
case "processing":
return "PROCESSING"
case "active":
return "ACTIVE"
case "failed":
return "FAILED"
default:
return status
}
}

func requiredAssetString(body map[string]any, key string) (string, error) {
value := strings.TrimSpace(stringValue(body, key))
if value == "" {
return "", fmt.Errorf("%s is required", key)
}
return value, nil
}

func stringValue(body map[string]any, key string) string {
if body == nil {
return ""
}
value, _ := body[key].(string)
return value
}

func mapValue(body map[string]any, key string) map[string]any {
if body == nil {
return nil
}
value, _ := body[key].(map[string]any)
return value
}

func intValue(body map[string]any, key string, fallback int) int {
if body == nil {
return fallback
}
switch value := body[key].(type) {
case int:
return value
case int32:
return int(value)
case int64:
return int(value)
case float64:
return int(value)
case float32:
return int(value)
default:
return fallback
}
}

func copyString(dst map[string]any, dstKey string, src map[string]any, srcKey string) {
value := stringValue(src, srcKey)
if strings.TrimSpace(value) != "" {
dst[dstKey] = value
}
}

func copyStringArray(dst map[string]any, dstKey string, src map[string]any, srcKey string) {
values := stringArrayValue(src, srcKey)
if len(values) > 0 {
dst[dstKey] = values
}
}

func copyStatusArray(dst map[string]any, dstKey string, src map[string]any, srcKey string) {
values := stringArrayValue(src, srcKey)
if len(values) == 0 {
return
}
mapped := make([]string, 0, len(values))
for _, value := range values {
mapped = append(mapped, officialStatusToChinaMobile(value))
}
dst[dstKey] = mapped
}

func stringArrayValue(body map[string]any, key string) []string {
if body == nil {
return nil
}
switch value := body[key].(type) {
case []string:
return value
case []any:
values := make([]string, 0, len(value))
for _, item := range value {
if s, ok := item.(string); ok && strings.TrimSpace(s) != "" {
values = append(values, s)
}
}
return values
default:
return nil
}
}

func compactAssetMap(input map[string]any) map[string]any {
output := make(map[string]any, len(input))
for key, value := range input {
switch v := value.(type) {
case string:
if strings.TrimSpace(v) != "" {
output[key] = v
}
case nil:
continue
default:
output[key] = value
}
}
return output
}

+ 41
- 0
service/asset_chinamobile_sdk.go Просмотреть файл

@@ -0,0 +1,41 @@
package service

import (
"gitlab.ecloud.com/ecloud/ecloudsdkcore/config"
"gitlab.ecloud.com/ecloud/ecloudsdkcore/utils"
"gitlab.ecloud.com/ecloud/ecloudsdkmaas"
cmmodel "gitlab.ecloud.com/ecloud/ecloudsdkmaas/model"
)

const (
chinaMobileAssetConnectTimeoutSeconds int32 = 10
chinaMobileAssetReadTimeoutSeconds int32 = 60
)

type chinaMobileAssetSDKClient interface {
CreateAsset(request *cmmodel.CreateAssetRequest) (*cmmodel.CreateAssetResponse, error)
ListAssets(request *cmmodel.ListAssetsRequest) (*cmmodel.ListAssetsResponse, error)
GetAsset(request *cmmodel.GetAssetRequest) (*cmmodel.GetAssetResponse, error)
UpdateAsset(request *cmmodel.UpdateAssetRequest) (*cmmodel.UpdateAssetResponse, error)
DeleteAsset(request *cmmodel.DeleteAssetRequest) (*cmmodel.DeleteAssetResponse, error)
CreateAssetGroup(request *cmmodel.CreateAssetGroupRequest) (*cmmodel.CreateAssetGroupResponse, error)
ListAssetGroups(request *cmmodel.ListAssetGroupsRequest) (*cmmodel.ListAssetGroupsResponse, error)
GetAssetGroup(request *cmmodel.GetAssetGroupRequest) (*cmmodel.GetAssetGroupResponse, error)
UpdateAssetGroup(request *cmmodel.UpdateAssetGroupRequest) (*cmmodel.UpdateAssetGroupResponse, error)
DeleteAssetGroup(request *cmmodel.DeleteAssetGroupRequest) (*cmmodel.DeleteAssetGroupResponse, error)
}

func newChinaMobileAssetSDKClient(credential chinaMobileAssetCredential) chinaMobileAssetSDKClient {
return ecloudsdkmaas.NewClient(newChinaMobileAssetSDKConfig(credential))
}

func newChinaMobileAssetSDKConfig(credential chinaMobileAssetCredential) *config.Config {
return &config.Config{
AccessKey: &credential.AK,
SecretKey: &credential.SK,
PoolId: &credential.PoolID,
ConnectTimeout: utils.Int32(chinaMobileAssetConnectTimeoutSeconds),
ReadTimeout: utils.Int32(chinaMobileAssetReadTimeoutSeconds),
IgnoreSSL: utils.Bool(false),
}
}

+ 19
- 0
service/asset_chinamobile_sdk_test.go Просмотреть файл

@@ -0,0 +1,19 @@
package service

import (
"testing"

"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)

func TestChinaMobileAssetSDKConfigUsesTimeoutsAndTLSVerification(t *testing.T) {
config := newChinaMobileAssetSDKConfig(chinaMobileAssetCredential{AK: "ak", SK: "sk", PoolID: "pool"})

require.NotNil(t, config.ConnectTimeout)
require.NotNil(t, config.ReadTimeout)
require.NotNil(t, config.IgnoreSSL)
assert.Greater(t, *config.ConnectTimeout, int32(0))
assert.Greater(t, *config.ReadTimeout, int32(0))
assert.False(t, *config.IgnoreSSL)
}

+ 467
- 0
service/asset_chinamobile_test.go Просмотреть файл

@@ -0,0 +1,467 @@
package service

import (
"context"
"errors"
"net/http"
"testing"
"time"

"github.com/QuantumNous/new-api/model"
"github.com/glebarez/sqlite"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
cmerrs "gitlab.ecloud.com/ecloud/ecloudsdkcore/errs"
cmmodel "gitlab.ecloud.com/ecloud/ecloudsdkmaas/model"
"gorm.io/gorm"
"gorm.io/gorm/logger"
)

func TestChinaMobileAssetAdapterCreateAssetMapsRequestAndResponse(t *testing.T) {
setChinaMobileAssetTestEnv(t, "")

fake := &fakeChinaMobileAssetSDKClient{}
adapter := &ChinaMobileAssetAdapter{newClient: func(credential chinaMobileAssetCredential) chinaMobileAssetSDKClient {
assert.Equal(t, "ak", credential.AK)
assert.Equal(t, "sk", credential.SK)
assert.Equal(t, defaultChinaMobileAssetPoolID, credential.PoolID)
return fake
}}

spec, ok := ParseAssetAction("CreateAsset")
require.True(t, ok)
resp, assetErr := adapter.DoAssetRequest(context.Background(), &model.Channel{Key: "video-generation-key"}, AssetRequest{
Action: spec,
Version: "2024-01-01",
Body: map[string]any{
"GroupId": "group-1",
"Name": "asset",
"URL": "https://example.com/a.png",
"AssetType": "Image",
},
})

require.Nil(t, assetErr)
require.NotNil(t, resp)
require.NotNil(t, fake.createAssetRequest)
assert.Equal(t, "group-1", *fake.createAssetRequest.CreateAssetBody.GroupId)
assert.Equal(t, "asset", *fake.createAssetRequest.CreateAssetBody.AssetName)
assert.Equal(t, "https://example.com/a.png", *fake.createAssetRequest.CreateAssetBody.AssetUrl)
assert.Equal(t, cmmodel.CreateAssetBodyAssetTypeEnumImage, *fake.createAssetRequest.CreateAssetBody.AssetType)
assert.Contains(t, string(resp.Body), `"RequestId":"req-1"`)
assert.Contains(t, string(resp.Body), `"Result":"asset-1"`)
}

func TestChinaMobileAssetCredentialFromChannelDefaultsCenterPool(t *testing.T) {
setChinaMobileAssetTestEnv(t, "")

credential, err := chinaMobileAssetCredentialFromChannel(&model.Channel{})

require.NoError(t, err)
assert.Equal(t, "ak", credential.AK)
assert.Equal(t, "sk", credential.SK)
assert.Equal(t, defaultChinaMobileAssetPoolID, credential.PoolID)
}

func TestChinaMobileAssetCredentialFromChannelSupportsPoolID(t *testing.T) {
setChinaMobileAssetTestEnv(t, "CIDC-RP-29")

credential, err := chinaMobileAssetCredentialFromChannel(&model.Channel{})

require.NoError(t, err)
assert.Equal(t, "ak", credential.AK)
assert.Equal(t, "sk", credential.SK)
assert.Equal(t, "CIDC-RP-29", credential.PoolID)
}

func TestChinaMobileAssetCredentialFromChannelRequiresCredential(t *testing.T) {
setupChinaMobileAssetTestDB(t)
t.Setenv(chinaMobileAssetAKEnv, "")
t.Setenv(chinaMobileAssetSKEnv, "")
t.Setenv(chinaMobileAssetPoolIDEnv, "")

_, err := chinaMobileAssetCredentialFromChannel(&model.Channel{Id: 99})

require.Error(t, err)
assert.Contains(t, err.Error(), "未配置素材凭证")
}

func TestChinaMobileAssetCredentialFromChannelFallsBackToLegacyEnvironment(t *testing.T) {
setupChinaMobileAssetTestDB(t)
t.Setenv(chinaMobileAssetAKEnv, "legacy-ak")
t.Setenv(chinaMobileAssetSKEnv, "legacy-sk")
t.Setenv(chinaMobileAssetPoolIDEnv, "legacy-pool")

credential, err := chinaMobileAssetCredentialFromChannel(&model.Channel{Id: 99})

require.NoError(t, err)
assert.Equal(t, "legacy-ak", credential.AK)
assert.Equal(t, "legacy-sk", credential.SK)
assert.Equal(t, "legacy-pool", credential.PoolID)
}

func TestChinaMobileAssetCredentialFromChannelPrefersChannelCredentialOverLegacyEnvironment(t *testing.T) {
setupChinaMobileAssetTestDB(t)
require.NoError(t, model.UpsertChannelAssetCredential(&model.ChannelAssetCredential{
ChannelId: 99,
AccessKey: "channel-ak",
SecretKey: "channel-sk",
PoolID: "channel-pool",
}))
t.Setenv(chinaMobileAssetAKEnv, "legacy-ak")
t.Setenv(chinaMobileAssetSKEnv, "legacy-sk")
t.Setenv(chinaMobileAssetPoolIDEnv, "legacy-pool")

credential, err := chinaMobileAssetCredentialFromChannel(&model.Channel{Id: 99})

require.NoError(t, err)
assert.Equal(t, "channel-ak", credential.AK)
assert.Equal(t, "channel-sk", credential.SK)
assert.Equal(t, "channel-pool", credential.PoolID)
}

func TestChinaMobileAssetAdapterReturnsBadRequestWhenCredentialIsMissing(t *testing.T) {
setupChinaMobileAssetTestDB(t)
adapter := &ChinaMobileAssetAdapter{newClient: func(credential chinaMobileAssetCredential) chinaMobileAssetSDKClient {
return &fakeChinaMobileAssetSDKClient{}
}}
spec, ok := ParseAssetAction("CreateAsset")
require.True(t, ok)

_, assetErr := adapter.DoAssetRequest(context.Background(), &model.Channel{Id: 61}, AssetRequest{Action: spec, Body: validChinaMobileCreateAssetBody()})

require.NotNil(t, assetErr)
assert.Equal(t, AssetErrorInvalidRequest, assetErr.Type)
assert.Equal(t, http.StatusBadRequest, assetErr.HTTPStatus)
assert.Contains(t, assetErr.Message, "未配置素材凭证")
}

func TestChinaMobileAssetAdapterAllOfficialActions(t *testing.T) {
cases := []struct {
action string
body map[string]any
assert func(t *testing.T, fake *fakeChinaMobileAssetSDKClient)
}{
{"CreateAssetGroup", map[string]any{"Name": "g", "GroupType": "AIGC"}, func(t *testing.T, fake *fakeChinaMobileAssetSDKClient) {
require.NotNil(t, fake.createAssetGroupRequest)
}},
{"CreateAsset", map[string]any{"GroupId": "g", "Name": "n", "URL": "https://example.com/a.png", "AssetType": "Image"}, func(t *testing.T, fake *fakeChinaMobileAssetSDKClient) { require.NotNil(t, fake.createAssetRequest) }},
{"ListAssetGroups", map[string]any{"Filter": map[string]any{"GroupType": "AIGC"}, "PageNumber": float64(1), "PageSize": float64(10)}, func(t *testing.T, fake *fakeChinaMobileAssetSDKClient) {
require.NotNil(t, fake.listAssetGroupsRequest)
}},
{"ListAssets", map[string]any{"Filter": map[string]any{"GroupType": "AIGC"}, "PageNumber": float64(1), "PageSize": float64(10)}, func(t *testing.T, fake *fakeChinaMobileAssetSDKClient) { require.NotNil(t, fake.listAssetsRequest) }},
{"GetAsset", map[string]any{"Id": "asset-1"}, func(t *testing.T, fake *fakeChinaMobileAssetSDKClient) {
require.NotNil(t, fake.getAssetRequest)
assert.Equal(t, "asset-1", *fake.getAssetRequest.GetAssetPath.AssetId)
}},
{"GetAssetGroup", map[string]any{"Id": "group-1"}, func(t *testing.T, fake *fakeChinaMobileAssetSDKClient) {
require.NotNil(t, fake.getAssetGroupRequest)
assert.Equal(t, "group-1", *fake.getAssetGroupRequest.GetAssetGroupPath.GroupId)
}},
{"UpdateAsset", map[string]any{"Id": "asset-1", "Name": "n"}, func(t *testing.T, fake *fakeChinaMobileAssetSDKClient) { require.NotNil(t, fake.updateAssetRequest) }},
{"UpdateAssetGroup", map[string]any{"Id": "group-1", "Name": "g"}, func(t *testing.T, fake *fakeChinaMobileAssetSDKClient) {
require.NotNil(t, fake.updateAssetGroupRequest)
}},
{"DeleteAsset", map[string]any{"Id": "asset-1"}, func(t *testing.T, fake *fakeChinaMobileAssetSDKClient) { require.NotNil(t, fake.deleteAssetRequest) }},
{"DeleteAssetGroup", map[string]any{"Id": "group-1"}, func(t *testing.T, fake *fakeChinaMobileAssetSDKClient) {
require.NotNil(t, fake.deleteAssetGroupRequest)
}},
}

for _, tc := range cases {
t.Run(tc.action, func(t *testing.T) {
setChinaMobileAssetTestEnv(t, "")

fake := &fakeChinaMobileAssetSDKClient{}
adapter := &ChinaMobileAssetAdapter{newClient: func(credential chinaMobileAssetCredential) chinaMobileAssetSDKClient {
return fake
}}
spec, ok := ParseAssetAction(tc.action)
require.True(t, ok)

resp, assetErr := adapter.DoAssetRequest(context.Background(), &model.Channel{Key: "video-generation-key"}, AssetRequest{
Action: spec,
Version: "2024-01-01",
Body: tc.body,
})

require.Nil(t, assetErr)
require.NotNil(t, resp)
tc.assert(t, fake)
assert.Contains(t, string(resp.Body), `"ResponseMetadata"`)
})
}
}

func TestChinaMobileAssetAdapterRejectsUnsupportedCompatibilityOptions(t *testing.T) {
setChinaMobileAssetTestEnv(t, "")

adapter := &ChinaMobileAssetAdapter{newClient: func(credential chinaMobileAssetCredential) chinaMobileAssetSDKClient {
return &fakeChinaMobileAssetSDKClient{}
}}
spec, ok := ParseAssetAction("ListAssets")
require.True(t, ok)

cases := []map[string]any{
{"ProjectName": "prod"},
{"SortBy": "Name"},
{"SortOrder": "Asc"},
}
for _, body := range cases {
_, assetErr := adapter.DoAssetRequest(context.Background(), &model.Channel{Key: "cm-key"}, AssetRequest{
Action: spec,
Body: body,
})

require.NotNil(t, assetErr)
assert.Equal(t, AssetErrorOperationNotSupported, assetErr.Type)
}
}

func TestChinaMobileAssetAdapterMissingIdReturnsInvalidRequest(t *testing.T) {
setChinaMobileAssetTestEnv(t, "")

adapter := &ChinaMobileAssetAdapter{newClient: func(credential chinaMobileAssetCredential) chinaMobileAssetSDKClient {
return &fakeChinaMobileAssetSDKClient{}
}}
spec, ok := ParseAssetAction("GetAsset")
require.True(t, ok)

_, assetErr := adapter.DoAssetRequest(context.Background(), &model.Channel{Key: "video-generation-key"}, AssetRequest{
Action: spec,
Body: map[string]any{},
})

require.NotNil(t, assetErr)
assert.Equal(t, AssetErrorInvalidRequest, assetErr.Type)
assert.Contains(t, assetErr.Message, "Id is required")
}

func TestChinaMobileAssetAdapterStateErrorMapsToUpstreamError(t *testing.T) {
_, assetErr := normalizeChinaMobileAssetResponse(AssetActionSpec{Action: "CreateAsset", Operation: AssetOperationAssetCreate}, "2024-01-01", []byte(`{"requestId":"req","state":"ERROR","errorCode":"Bad","errorMessage":"bad request"}`))

require.NotNil(t, assetErr)
assert.Equal(t, AssetErrorUpstream, assetErr.Type)
assert.Equal(t, "bad request", assetErr.Message)
}

func TestChinaMobileAssetAdapterRequiresGroupTypeForListAssets(t *testing.T) {
setChinaMobileAssetTestEnv(t, "")

adapter := &ChinaMobileAssetAdapter{newClient: func(credential chinaMobileAssetCredential) chinaMobileAssetSDKClient {
return &fakeChinaMobileAssetSDKClient{}
}}
spec, ok := ParseAssetAction("ListAssets")
require.True(t, ok)

_, assetErr := adapter.DoAssetRequest(context.Background(), &model.Channel{Key: "video-generation-key"}, AssetRequest{
Action: spec,
Body: map[string]any{"PageNumber": 1, "PageSize": 10},
})

require.NotNil(t, assetErr)
assert.Equal(t, AssetErrorInvalidRequest, assetErr.Type)
assert.Equal(t, http.StatusBadRequest, assetErr.HTTPStatus)
}

func TestChinaMobileAssetDeleteFalseMapsToUpstreamError(t *testing.T) {
_, assetErr := normalizeChinaMobileAssetResponse(AssetActionSpec{Action: "DeleteAsset", Operation: AssetOperationAssetDelete, Delete: true}, "2024-01-01", []byte(`{"requestId":"req","state":"OK","body":false}`))

require.NotNil(t, assetErr)
assert.Equal(t, AssetErrorUpstream, assetErr.Type)
}

func TestChinaMobileAssetUnknownStateMapsToUpstreamError(t *testing.T) {
_, assetErr := normalizeChinaMobileAssetResponse(AssetActionSpec{Action: "CreateAsset", Operation: AssetOperationAssetCreate}, "2024-01-01", []byte(`{"requestId":"req","state":"","body":"asset-1"}`))

require.NotNil(t, assetErr)
assert.Equal(t, AssetErrorUpstream, assetErr.Type)
}

func TestChinaMobileAssetSDKBadRequestMapsToInvalidRequest(t *testing.T) {
setChinaMobileAssetTestEnv(t, "")

adapter := &ChinaMobileAssetAdapter{newClient: func(credential chinaMobileAssetCredential) chinaMobileAssetSDKClient {
return &errorChinaMobileAssetSDKClient{
fakeChinaMobileAssetSDKClient: fakeChinaMobileAssetSDKClient{},
err: cmerrs.NewServerResponseError("bad request", nil, http.StatusBadRequest, nil, `{"errorMessage":"invalid asset"}`),
}
}}
spec, ok := ParseAssetAction("CreateAsset")
require.True(t, ok)

_, assetErr := adapter.DoAssetRequest(context.Background(), &model.Channel{Key: "video-generation-key"}, AssetRequest{Action: spec, Body: validChinaMobileCreateAssetBody()})

require.NotNil(t, assetErr)
assert.Equal(t, AssetErrorInvalidRequest, assetErr.Type)
assert.Equal(t, http.StatusBadRequest, assetErr.HTTPStatus)
}

func TestChinaMobileAssetSDKServerErrorRemainsUpstreamError(t *testing.T) {
setChinaMobileAssetTestEnv(t, "")

adapter := &ChinaMobileAssetAdapter{newClient: func(credential chinaMobileAssetCredential) chinaMobileAssetSDKClient {
return &errorChinaMobileAssetSDKClient{
fakeChinaMobileAssetSDKClient: fakeChinaMobileAssetSDKClient{},
err: cmerrs.NewServerResponseError("server error", nil, http.StatusInternalServerError, nil, ""),
}
}}
spec, ok := ParseAssetAction("CreateAsset")
require.True(t, ok)

_, assetErr := adapter.DoAssetRequest(context.Background(), &model.Channel{Key: "video-generation-key"}, AssetRequest{Action: spec, Body: validChinaMobileCreateAssetBody()})

require.NotNil(t, assetErr)
assert.Equal(t, AssetErrorUpstream, assetErr.Type)
assert.Equal(t, http.StatusBadGateway, assetErr.HTTPStatus)
}

func TestCallChinaMobileAssetSDKReturnsWhenContextCancelled(t *testing.T) {
client := &blockingChinaMobileAssetSDKClient{fakeChinaMobileAssetSDKClient: fakeChinaMobileAssetSDKClient{}, release: make(chan struct{})}
ctx, cancel := context.WithCancel(context.Background())
cancel()
spec, ok := ParseAssetAction("CreateAsset")
require.True(t, ok)

started := time.Now()
_, err := callChinaMobileAssetSDK(ctx, client, AssetRequest{Action: spec, Body: validChinaMobileCreateAssetBody()})

require.ErrorIs(t, err, context.Canceled)
assert.Less(t, time.Since(started), 500*time.Millisecond)
close(client.release)
}

func setChinaMobileAssetTestEnv(t *testing.T, poolID string) {
t.Helper()
setupChinaMobileAssetTestDB(t)
require.NoError(t, model.UpsertChannelAssetCredential(&model.ChannelAssetCredential{
ChannelId: 0,
AccessKey: "ak",
SecretKey: "sk",
PoolID: poolID,
}))
}

func setupChinaMobileAssetTestDB(t *testing.T) *gorm.DB {
t.Helper()
db, err := gorm.Open(sqlite.Open("file:service_chinamobile_asset?mode=memory&cache=shared"), &gorm.Config{
Logger: logger.Default.LogMode(logger.Silent),
})
require.NoError(t, err)
sqlDB, err := db.DB()
require.NoError(t, err)
sqlDB.SetMaxOpenConns(1)
originalDB := model.DB
model.DB = db
require.NoError(t, db.AutoMigrate(&model.ChannelAssetCredential{}))
t.Cleanup(func() {
model.DB = originalDB
require.NoError(t, sqlDB.Close())
})
return db
}

type blockingChinaMobileAssetSDKClient struct {
fakeChinaMobileAssetSDKClient
release chan struct{}
}

type errorChinaMobileAssetSDKClient struct {
fakeChinaMobileAssetSDKClient
err error
}

func (f *errorChinaMobileAssetSDKClient) CreateAsset(request *cmmodel.CreateAssetRequest) (*cmmodel.CreateAssetResponse, error) {
if f.err == nil {
return nil, errors.New("missing test error")
}
return nil, f.err
}

func (f *blockingChinaMobileAssetSDKClient) CreateAsset(request *cmmodel.CreateAssetRequest) (*cmmodel.CreateAssetResponse, error) {
<-f.release
return f.fakeChinaMobileAssetSDKClient.CreateAsset(request)
}

func validChinaMobileCreateAssetBody() map[string]any {
return map[string]any{
"GroupId": "group-1",
"Name": "asset",
"URL": "https://example.com/asset.png",
"AssetType": "Image",
}
}

type fakeChinaMobileAssetSDKClient struct {
createAssetRequest *cmmodel.CreateAssetRequest
listAssetsRequest *cmmodel.ListAssetsRequest
getAssetRequest *cmmodel.GetAssetRequest
updateAssetRequest *cmmodel.UpdateAssetRequest
deleteAssetRequest *cmmodel.DeleteAssetRequest
createAssetGroupRequest *cmmodel.CreateAssetGroupRequest
listAssetGroupsRequest *cmmodel.ListAssetGroupsRequest
getAssetGroupRequest *cmmodel.GetAssetGroupRequest
updateAssetGroupRequest *cmmodel.UpdateAssetGroupRequest
deleteAssetGroupRequest *cmmodel.DeleteAssetGroupRequest
}

func (f *fakeChinaMobileAssetSDKClient) CreateAsset(request *cmmodel.CreateAssetRequest) (*cmmodel.CreateAssetResponse, error) {
f.createAssetRequest = request
return (&cmmodel.CreateAssetResponse{}).SetRequestId("req-1").SetState("OK").SetBody("asset-1"), nil
}

func (f *fakeChinaMobileAssetSDKClient) ListAssets(request *cmmodel.ListAssetsRequest) (*cmmodel.ListAssetsResponse, error) {
f.listAssetsRequest = request
item := cmmodel.ListAssetsResponseData{}
item.SetAssetId("asset-1").SetGroupId("group-1").SetAssetName("n").SetStatus(cmmodel.ListAssetsResponseDataStatusEnumActive)
body := (&cmmodel.ListAssetsResponseBody{}).SetData([]cmmodel.ListAssetsResponseData{item}).SetTotal(1)
return (&cmmodel.ListAssetsResponse{}).SetRequestId("req").SetState("OK").SetBody(body), nil
}

func (f *fakeChinaMobileAssetSDKClient) GetAsset(request *cmmodel.GetAssetRequest) (*cmmodel.GetAssetResponse, error) {
f.getAssetRequest = request
body := (&cmmodel.GetAssetResponseBody{}).SetAssetId("asset-1").SetGroupId("group-1").SetAssetName("n").SetStatus(cmmodel.GetAssetResponseBodyStatusEnumActive)
return (&cmmodel.GetAssetResponse{}).SetRequestId("req").SetState("OK").SetBody(body), nil
}

func (f *fakeChinaMobileAssetSDKClient) UpdateAsset(request *cmmodel.UpdateAssetRequest) (*cmmodel.UpdateAssetResponse, error) {
f.updateAssetRequest = request
body := (&cmmodel.UpdateAssetResponseBody{}).SetAssetId("asset-1").SetGroupId("group-1").SetAssetName("n").SetStatus(cmmodel.UpdateAssetResponseBodyStatusEnumActive)
return (&cmmodel.UpdateAssetResponse{}).SetRequestId("req").SetState("OK").SetBody(body), nil
}

func (f *fakeChinaMobileAssetSDKClient) DeleteAsset(request *cmmodel.DeleteAssetRequest) (*cmmodel.DeleteAssetResponse, error) {
f.deleteAssetRequest = request
return (&cmmodel.DeleteAssetResponse{}).SetRequestId("req").SetState("OK").SetBody(true), nil
}

func (f *fakeChinaMobileAssetSDKClient) CreateAssetGroup(request *cmmodel.CreateAssetGroupRequest) (*cmmodel.CreateAssetGroupResponse, error) {
f.createAssetGroupRequest = request
body := (&cmmodel.CreateAssetGroupResponseBody{}).SetGroupId("group-1").SetGroupName("g").SetGroupType(cmmodel.CreateAssetGroupResponseBodyGroupTypeEnumAigc)
return (&cmmodel.CreateAssetGroupResponse{}).SetRequestId("req").SetState("OK").SetBody(body), nil
}

func (f *fakeChinaMobileAssetSDKClient) ListAssetGroups(request *cmmodel.ListAssetGroupsRequest) (*cmmodel.ListAssetGroupsResponse, error) {
f.listAssetGroupsRequest = request
item := cmmodel.ListAssetGroupsResponseData{}
item.SetGroupId("group-1").SetGroupName("g").SetGroupType(cmmodel.ListAssetGroupsResponseDataGroupTypeEnumAigc)
body := (&cmmodel.ListAssetGroupsResponseBody{}).SetData([]cmmodel.ListAssetGroupsResponseData{item}).SetTotal(1)
return (&cmmodel.ListAssetGroupsResponse{}).SetRequestId("req").SetState("OK").SetBody(body), nil
}

func (f *fakeChinaMobileAssetSDKClient) GetAssetGroup(request *cmmodel.GetAssetGroupRequest) (*cmmodel.GetAssetGroupResponse, error) {
f.getAssetGroupRequest = request
body := (&cmmodel.GetAssetGroupResponseBody{}).SetGroupId("group-1").SetGroupName("g").SetGroupType(cmmodel.GetAssetGroupResponseBodyGroupTypeEnumAigc)
return (&cmmodel.GetAssetGroupResponse{}).SetRequestId("req").SetState("OK").SetBody(body), nil
}

func (f *fakeChinaMobileAssetSDKClient) UpdateAssetGroup(request *cmmodel.UpdateAssetGroupRequest) (*cmmodel.UpdateAssetGroupResponse, error) {
f.updateAssetGroupRequest = request
body := (&cmmodel.UpdateAssetGroupResponseBody{}).SetGroupId("group-1").SetGroupName("g").SetGroupType(cmmodel.UpdateAssetGroupResponseBodyGroupTypeEnumAigc)
return (&cmmodel.UpdateAssetGroupResponse{}).SetRequestId("req").SetState("OK").SetBody(body), nil
}

func (f *fakeChinaMobileAssetSDKClient) DeleteAssetGroup(request *cmmodel.DeleteAssetGroupRequest) (*cmmodel.DeleteAssetGroupResponse, error) {
f.deleteAssetGroupRequest = request
return (&cmmodel.DeleteAssetGroupResponse{}).SetRequestId("req").SetState("OK").SetBody(true), nil
}

+ 98
- 0
service/asset_compatible.go Просмотреть файл

@@ -0,0 +1,98 @@
package service

import (
"bytes"
"context"
"fmt"
"io"
"net/http"
"net/url"
"strings"

"github.com/QuantumNous/new-api/common"
"github.com/QuantumNous/new-api/model"
"github.com/QuantumNous/new-api/setting/system_setting"
)

type CompatibleAssetAdapter struct {
name string
operation map[AssetOperation]struct{}
}

func NewCompatibleAssetAdapter(name string, operations []AssetOperation) AssetAdapter {
supported := make(map[AssetOperation]struct{}, len(operations))
for _, op := range operations {
supported[op] = struct{}{}
}
return &CompatibleAssetAdapter{name: name, operation: supported}
}

func (a *CompatibleAssetAdapter) Name() string {
return a.name
}

func (a *CompatibleAssetAdapter) Supports(operation AssetOperation) bool {
_, ok := a.operation[operation]
return ok
}

func (a *CompatibleAssetAdapter) DoAssetRequest(ctx context.Context, channel *model.Channel, req AssetRequest) (*AssetUpstreamResponse, *AssetError) {
if !a.Supports(req.Action.Operation) {
return nil, newAssetError(AssetErrorOperationNotSupported, fmt.Sprintf("asset operation %s is not supported", req.Action.Operation), http.StatusBadRequest)
}
upstreamURL, err := buildCompatibleAssetURL(channel, req.Action.Action, req.Version)
if err != nil {
return nil, newAssetError(AssetErrorServer, err.Error(), http.StatusInternalServerError)
}
fetchSetting := system_setting.GetFetchSetting()
if err := common.ValidateURLWithFetchSetting(upstreamURL, fetchSetting.EnableSSRFProtection, fetchSetting.AllowPrivateIp, fetchSetting.DomainFilterMode, fetchSetting.IpFilterMode, fetchSetting.DomainList, fetchSetting.IpList, fetchSetting.AllowedPorts, fetchSetting.ApplyIPFilterForDomain); err != nil {
return nil, newAssetError(AssetErrorServer, fmt.Sprintf("request blocked: %v", err), http.StatusForbidden)
}
httpReq, err := http.NewRequestWithContext(ctx, http.MethodPost, upstreamURL, bytes.NewReader(req.RawBody))
if err != nil {
return nil, newAssetError(AssetErrorServer, err.Error(), http.StatusInternalServerError)
}
httpReq.Header.Set("Content-Type", "application/json")
httpReq.Header.Set("Accept", "application/json")
httpReq.Header.Set("Authorization", "Bearer "+strings.TrimSpace(channel.Key))
client, err := GetHttpClientWithProxy(channel.GetSetting().Proxy)
if err != nil {
return nil, newAssetError(AssetErrorServer, err.Error(), http.StatusInternalServerError)
}
if client == nil {
client = http.DefaultClient
}
resp, err := client.Do(httpReq)
if err != nil {
return nil, newAssetError(AssetErrorUpstream, err.Error(), http.StatusBadGateway)
}
defer resp.Body.Close()
data, err := io.ReadAll(resp.Body)
if err != nil {
return nil, newAssetError(AssetErrorUpstream, err.Error(), http.StatusBadGateway)
}
if resp.StatusCode < 200 || resp.StatusCode >= 300 {
return nil, newAssetError(AssetErrorUpstream, string(data), http.StatusBadGateway)
}
return &AssetUpstreamResponse{StatusCode: resp.StatusCode, Header: resp.Header, Body: data}, nil
}

func buildCompatibleAssetURL(channel *model.Channel, action string, version string) (string, error) {
baseURL := strings.TrimSpace(channel.GetBaseURL())
if baseURL == "" {
baseURL = "https://ark.cn-beijing.volcengineapi.com"
}
u, err := url.Parse(baseURL)
if err != nil {
return "", err
}
u.Path = strings.TrimRight(u.Path, "/") + "/api/v1/multimodal/sd/assets"
q := u.Query()
q.Set("Action", action)
if strings.TrimSpace(version) == "" {
version = "2024-01-01"
}
q.Set("Version", version)
u.RawQuery = q.Encode()
return u.String(), nil
}

+ 27
- 0
service/asset_compatible_test.go Просмотреть файл

@@ -0,0 +1,27 @@
package service

import (
"context"
"net/http"
"testing"

"github.com/QuantumNous/new-api/model"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)

func TestCompatibleAssetAdapterBlocksPrivateBaseURL(t *testing.T) {
baseURL := "http://127.0.0.1:8080"
adapter := NewCompatibleAssetAdapter("compatible", []AssetOperation{AssetOperationAssetCreate})
spec, ok := ParseAssetAction("CreateAsset")
require.True(t, ok)

_, assetErr := adapter.DoAssetRequest(context.Background(), &model.Channel{Key: "key", BaseURL: &baseURL}, AssetRequest{
Action: spec,
Version: "2024-01-01",
RawBody: []byte(`{}`),
})

require.NotNil(t, assetErr)
assert.Equal(t, http.StatusForbidden, assetErr.HTTPStatus)
}

+ 77
- 0
service/asset_operation.go Просмотреть файл

@@ -0,0 +1,77 @@
package service

import (
"strings"

"github.com/QuantumNous/new-api/common"
)

type AssetOperation string

const (
AssetOperationAssetCreate AssetOperation = "asset.create"
AssetOperationAssetList AssetOperation = "asset.list"
AssetOperationAssetGet AssetOperation = "asset.get"
AssetOperationAssetUpdate AssetOperation = "asset.update"
AssetOperationAssetDelete AssetOperation = "asset.delete"
AssetOperationAssetGroupCreate AssetOperation = "asset_group.create"
AssetOperationAssetGroupList AssetOperation = "asset_group.list"
AssetOperationAssetGroupGet AssetOperation = "asset_group.get"
AssetOperationAssetGroupUpdate AssetOperation = "asset_group.update"
AssetOperationAssetGroupDelete AssetOperation = "asset_group.delete"
)

type AssetActionSpec struct {
Action string
Operation AssetOperation
Delete bool
}

var assetActionSpecs = map[string]AssetActionSpec{
"createassetgroup": {Action: "CreateAssetGroup", Operation: AssetOperationAssetGroupCreate},
"createasset": {Action: "CreateAsset", Operation: AssetOperationAssetCreate},
"listassetgroups": {Action: "ListAssetGroups", Operation: AssetOperationAssetGroupList},
"listassets": {Action: "ListAssets", Operation: AssetOperationAssetList},
"getasset": {Action: "GetAsset", Operation: AssetOperationAssetGet},
"getassetgroup": {Action: "GetAssetGroup", Operation: AssetOperationAssetGroupGet},
"updateasset": {Action: "UpdateAsset", Operation: AssetOperationAssetUpdate},
"updateassetgroup": {Action: "UpdateAssetGroup", Operation: AssetOperationAssetGroupUpdate},
"deleteasset": {Action: "DeleteAsset", Operation: AssetOperationAssetDelete, Delete: true},
"deleteassetgroup": {Action: "DeleteAssetGroup", Operation: AssetOperationAssetGroupDelete, Delete: true},
}

func ParseAssetAction(action string) (AssetActionSpec, bool) {
spec, ok := assetActionSpecs[strings.ToLower(strings.TrimSpace(action))]
return spec, ok
}

type assetResponseMetadata struct {
RequestID string `json:"RequestId,omitempty"`
Action string `json:"Action"`
Version string `json:"Version"`
Service string `json:"Service"`
Region string `json:"Region"`
}

func BuildAssetSuccessResponse(action string, version string, result any) ([]byte, error) {
return buildAssetSuccessResponseWithRequestID(action, version, "", result)
}

func buildAssetSuccessResponseWithRequestID(action string, version string, requestID string, result any) ([]byte, error) {
if strings.TrimSpace(version) == "" {
version = "2024-01-01"
}
payload := map[string]any{
"ResponseMetadata": assetResponseMetadata{
RequestID: requestID,
Action: action,
Version: version,
Service: "ark",
Region: "cn-beijing",
},
}
if result != nil {
payload["Result"] = result
}
return common.Marshal(payload)
}

+ 45
- 0
service/asset_operation_test.go Просмотреть файл

@@ -0,0 +1,45 @@
package service

import (
"testing"

"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)

func TestParseAssetActionRecognizesOfficialActions(t *testing.T) {
cases := map[string]AssetOperation{
"CreateAssetGroup": AssetOperationAssetGroupCreate,
"CreateAsset": AssetOperationAssetCreate,
"ListAssetGroups": AssetOperationAssetGroupList,
"ListAssets": AssetOperationAssetList,
"GetAsset": AssetOperationAssetGet,
"GetAssetGroup": AssetOperationAssetGroupGet,
"UpdateAsset": AssetOperationAssetUpdate,
"UpdateAssetGroup": AssetOperationAssetGroupUpdate,
"DeleteAsset": AssetOperationAssetDelete,
"DeleteAssetGroup": AssetOperationAssetGroupDelete,
}

for action, operation := range cases {
t.Run(action, func(t *testing.T) {
spec, ok := ParseAssetAction(action)
require.True(t, ok)
assert.Equal(t, action, spec.Action)
assert.Equal(t, operation, spec.Operation)
})
}
}

func TestParseAssetActionRejectsUnknownAction(t *testing.T) {
_, ok := ParseAssetAction("CreateRealPersonAuthSession")
assert.False(t, ok)
}

func TestBuildAssetSuccessResponseAllowsDeleteWithoutResult(t *testing.T) {
body, err := BuildAssetSuccessResponse("DeleteAsset", "2024-01-01", nil)
require.NoError(t, err)
assert.Contains(t, string(body), `"Action":"DeleteAsset"`)
assert.Contains(t, string(body), `"ResponseMetadata"`)
assert.NotContains(t, string(body), `"Result"`)
}

+ 79
- 0
service/asset_resolver.go Просмотреть файл

@@ -0,0 +1,79 @@
package service

import (
"errors"
"fmt"
"net/http"

"github.com/QuantumNous/new-api/common"
"github.com/QuantumNous/new-api/model"
"gorm.io/gorm"
)

func ResolveAssetChannelForOperation(userID int, tokenGroup string, operation AssetOperation) (*model.Channel, AssetAdapter, *AssetError) {
bindings, err := model.GetUserAssetChannelsByTypes(userID, RegisteredAssetChannelTypes(), tokenGroup)
if err != nil {
return nil, nil, newAssetError(AssetErrorServer, err.Error(), http.StatusInternalServerError)
}
if len(bindings) > 0 {
binding := bindings[0]
channel, err := model.CacheGetChannel(binding.ChannelId)
if err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) || common.MemoryCacheEnabled {
return nil, nil, newAssetError(AssetErrorBindingInvalid, "bound asset channel not found", http.StatusBadGateway)
}
return nil, nil, newAssetError(AssetErrorServer, err.Error(), http.StatusInternalServerError)
}
adapter, ok := GetAssetAdapter(channel.Type)
if channel.Status != common.ChannelStatusEnabled || !assetChannelHasKey(channel) || !MatchDoubaoAssetGroup(channel.GetGroups(), tokenGroup) || !ok {
return nil, nil, newAssetError(AssetErrorBindingInvalid, "bound asset channel is not available for asset library", http.StatusBadGateway)
}
if !adapter.Supports(operation) {
return nil, nil, newAssetError(AssetErrorOperationNotSupported, fmt.Sprintf("asset operation %s is not supported by bound channel", operation), http.StatusBadRequest)
}
return channel, adapter, nil
}

channel, adapter, err := autoMatchAssetChannelForOperation(tokenGroup, operation)
if err != nil {
return nil, nil, newAssetError(AssetErrorServer, err.Error(), http.StatusInternalServerError)
}
if channel == nil {
return nil, nil, newAssetError(AssetErrorChannelNotFound, "no available asset channel supports requested operation", http.StatusBadGateway)
}
return channel, adapter, nil
}

func autoMatchAssetChannelForOperation(tokenGroup string, operation AssetOperation) (*model.Channel, AssetAdapter, error) {
for _, channelType := range RegisteredAssetChannelTypes() {
adapter, ok := GetAssetAdapter(channelType)
if !ok || !adapter.Supports(operation) {
continue
}
for startIdx := 0; ; startIdx += DoubaoAssetChannelPageSize {
candidates, err := model.GetChannelsByType(startIdx, DoubaoAssetChannelPageSize, true, channelType)
if err != nil {
return nil, nil, err
}
for _, candidate := range candidates {
if candidate == nil || !MatchDoubaoAssetGroup(candidate.GetGroups(), tokenGroup) {
continue
}
channel, err := model.CacheGetChannel(candidate.Id)
if err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) || common.MemoryCacheEnabled {
continue
}
return nil, nil, err
}
if channel.Status == common.ChannelStatusEnabled && assetChannelHasKey(channel) && MatchDoubaoAssetGroup(channel.GetGroups(), tokenGroup) {
return channel, adapter, nil
}
}
if len(candidates) < DoubaoAssetChannelPageSize {
break
}
}
}
return nil, nil, nil
}

+ 106
- 0
service/asset_resolver_test.go Просмотреть файл

@@ -0,0 +1,106 @@
package service

import (
"context"
"net/http"
"testing"

"github.com/QuantumNous/new-api/common"
"github.com/QuantumNous/new-api/constant"
"github.com/QuantumNous/new-api/model"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)

type fakeAssetAdapter struct {
name string
operation AssetOperation
statusCode int
}

func (a fakeAssetAdapter) Name() string {
return a.name
}

func (a fakeAssetAdapter) Supports(op AssetOperation) bool {
return a.operation == op
}

func (a fakeAssetAdapter) DoAssetRequest(context.Context, *model.Channel, AssetRequest) (*AssetUpstreamResponse, *AssetError) {
return &AssetUpstreamResponse{StatusCode: a.statusCode, Body: []byte(`{"ok":true}`)}, nil
}

func TestResolveAssetChannelAutoMatchesOnlyRegisteredOperation(t *testing.T) {
db := setupDoubaoAssetChannelDB(t)
resetAssetAdapterRegistryForTest(t)
RegisterAssetAdapterForTest(constant.ChannelTypeChinaMobileSeedance, fakeAssetAdapter{name: "cm", operation: AssetOperationAssetCreate, statusCode: http.StatusOK})
createDoubaoAssetChannelForTest(t, db, 1, constant.ChannelTypeOpenAI, "default", "openai-key", common.ChannelStatusEnabled)
createDoubaoAssetChannelForTest(t, db, 2, constant.ChannelTypeChinaMobileSeedance, "default", "cm-key", common.ChannelStatusEnabled)

ch, adapter, assetErr := ResolveAssetChannelForOperation(10, "default", AssetOperationAssetCreate)

require.Nil(t, assetErr)
require.NotNil(t, ch)
require.NotNil(t, adapter)
assert.Equal(t, 2, ch.Id)
assert.Equal(t, "cm", adapter.Name())
}

func TestResolveAssetChannelBoundUnsupportedOperationDoesNotFallback(t *testing.T) {
db := setupDoubaoAssetChannelDB(t)
resetAssetAdapterRegistryForTest(t)
RegisterAssetAdapterForTest(constant.ChannelTypeChinaMobileSeedance, fakeAssetAdapter{name: "cm", operation: AssetOperationAssetCreate})
RegisterAssetAdapterForTest(constant.ChannelTypeDoubaoVideoCompatibleAiping, fakeAssetAdapter{name: "aiping", operation: AssetOperationAssetGroupCreate})
createDoubaoAssetChannelForTest(t, db, 1, constant.ChannelTypeChinaMobileSeedance, "default", "cm-key", common.ChannelStatusEnabled)
createDoubaoAssetChannelForTest(t, db, 2, constant.ChannelTypeDoubaoVideoCompatibleAiping, "default", "aiping-key", common.ChannelStatusEnabled)
require.NoError(t, model.BindUserAssetChannel(10, constant.ChannelTypeChinaMobileSeedance, "default", 1))

ch, adapter, assetErr := ResolveAssetChannelForOperation(10, "default", AssetOperationAssetGroupCreate)

assert.Nil(t, ch)
assert.Nil(t, adapter)
require.NotNil(t, assetErr)
assert.Equal(t, AssetErrorOperationNotSupported, assetErr.Type)
}

func TestResolveAssetChannelIgnoresBindingFromUnrelatedFamily(t *testing.T) {
db := setupDoubaoAssetChannelDB(t)
resetAssetAdapterRegistryForTest(t)
RegisterAssetAdapterForTest(constant.ChannelTypeChinaMobileSeedance, fakeAssetAdapter{name: "cm", operation: AssetOperationAssetCreate})
createDoubaoAssetChannelForTest(t, db, 1, constant.ChannelTypeKlingAiping, "default", "kling-key", common.ChannelStatusEnabled)
createDoubaoAssetChannelForTest(t, db, 2, constant.ChannelTypeChinaMobileSeedance, "default", "cm-key", common.ChannelStatusEnabled)
require.NoError(t, model.BindUserAssetChannel(10, constant.ChannelTypeKlingAiping, "default", 1))

ch, adapter, assetErr := ResolveAssetChannelForOperation(10, "default", AssetOperationAssetCreate)

require.Nil(t, assetErr)
require.NotNil(t, ch)
require.NotNil(t, adapter)
assert.Equal(t, 2, ch.Id)
}

func TestResolveAssetChannelAutoMatchDoesNotReplaceVideoBinding(t *testing.T) {
db := setupDoubaoAssetChannelDB(t)
resetAssetAdapterRegistryForTest(t)
RegisterAssetAdapterForTest(constant.ChannelTypeChinaMobileSeedance, fakeAssetAdapter{name: "cm", operation: AssetOperationAssetGroupCreate})
createDoubaoAssetChannelForTest(t, db, 1, constant.ChannelTypeKlingAiping, "default", "kling-key", common.ChannelStatusEnabled)
createDoubaoAssetChannelForTest(t, db, 2, constant.ChannelTypeChinaMobileSeedance, "default", "cm-key", common.ChannelStatusEnabled)
require.NoError(t, model.BindUserAssetChannel(10, constant.ChannelTypeKlingAiping, "default", 1))

ch, _, assetErr := ResolveAssetChannelForOperation(10, "default", AssetOperationAssetGroupCreate)

require.Nil(t, assetErr)
require.NotNil(t, ch)
assert.Equal(t, 2, ch.Id)
binding, err := model.GetUserAssetChannel(10, constant.ChannelTypeKlingAiping, "default")
require.NoError(t, err)
require.NotNil(t, binding)
assert.Equal(t, 1, binding.ChannelId)
}

func TestCompatibleAssetAdapterSupportsOnlyDeclaredOperations(t *testing.T) {
adapter := NewCompatibleAssetAdapter("aiping_asset", []AssetOperation{AssetOperationAssetCreate})

assert.True(t, adapter.Supports(AssetOperationAssetCreate))
assert.False(t, adapter.Supports(AssetOperationAssetGroupCreate))
}

+ 163
- 0
service/asset_tianyiyun.go Просмотреть файл

@@ -0,0 +1,163 @@
package service

import (
"bytes"
"context"
"fmt"
"io"
"net/http"
"net/url"
"strings"

"github.com/QuantumNous/new-api/common"
"github.com/QuantumNous/new-api/model"
"github.com/QuantumNous/new-api/setting/system_setting"
)

type TianyiYunAssetAdapter struct{}

func NewTianyiYunAssetAdapter() AssetAdapter {
return &TianyiYunAssetAdapter{}
}

func (a *TianyiYunAssetAdapter) Name() string {
return "tianyiyun_asset"
}

func (a *TianyiYunAssetAdapter) Supports(operation AssetOperation) bool {
return operation == AssetOperationAssetCreate || operation == AssetOperationAssetGet
}

func (a *TianyiYunAssetAdapter) DoAssetRequest(ctx context.Context, channel *model.Channel, req AssetRequest) (*AssetUpstreamResponse, *AssetError) {
if !a.Supports(req.Action.Operation) {
return nil, newAssetError(AssetErrorOperationNotSupported, fmt.Sprintf("asset operation %s is not supported", req.Action.Operation), http.StatusBadRequest)
}
if channel == nil {
return nil, newAssetError(AssetErrorInvalidRequest, "channel is required", http.StatusBadRequest)
}

upstreamURL, body, method, assetErr := buildTianyiYunAssetRequest(channel, req)
if assetErr != nil {
return nil, assetErr
}
fetchSetting := system_setting.GetFetchSetting()
if err := common.ValidateURLWithFetchSetting(upstreamURL, fetchSetting.EnableSSRFProtection, fetchSetting.AllowPrivateIp, fetchSetting.DomainFilterMode, fetchSetting.IpFilterMode, fetchSetting.DomainList, fetchSetting.IpList, fetchSetting.AllowedPorts, fetchSetting.ApplyIPFilterForDomain); err != nil {
return nil, newAssetError(AssetErrorServer, fmt.Sprintf("request blocked: %v", err), http.StatusForbidden)
}

httpReq, err := http.NewRequestWithContext(ctx, method, upstreamURL, bytes.NewReader(body))
if err != nil {
return nil, newAssetError(AssetErrorServer, err.Error(), http.StatusInternalServerError)
}
httpReq.Header.Set("Accept", "application/json")
httpReq.Header.Set("Authorization", "Bearer "+strings.TrimSpace(channel.Key))
if method == http.MethodPost {
httpReq.Header.Set("Content-Type", "application/json")
}
client, err := GetHttpClientWithProxy(channel.GetSetting().Proxy)
if err != nil {
return nil, newAssetError(AssetErrorServer, err.Error(), http.StatusInternalServerError)
}
if client == nil {
client = http.DefaultClient
}
resp, err := client.Do(httpReq)
if err != nil {
return nil, newAssetError(AssetErrorUpstream, err.Error(), http.StatusBadGateway)
}
defer resp.Body.Close()
data, err := io.ReadAll(resp.Body)
if err != nil {
return nil, newAssetError(AssetErrorUpstream, err.Error(), http.StatusBadGateway)
}
if resp.StatusCode < http.StatusOK || resp.StatusCode >= http.StatusMultipleChoices {
return nil, newAssetError(AssetErrorUpstream, string(data), http.StatusBadGateway)
}
return &AssetUpstreamResponse{StatusCode: resp.StatusCode, Header: resp.Header, Body: data}, nil
}

func buildTianyiYunAssetRequest(channel *model.Channel, req AssetRequest) (string, []byte, string, *AssetError) {
baseURL, err := tianyiYunAssetBaseURL(channel)
if err != nil {
return "", nil, "", newAssetError(AssetErrorInvalidRequest, err.Error(), http.StatusBadRequest)
}
switch req.Action.Operation {
case AssetOperationAssetCreate:
payload, assetErr := buildTianyiYunAssetCreatePayload(req.Body)
if assetErr != nil {
return "", nil, "", assetErr
}
data, err := common.Marshal(payload)
if err != nil {
return "", nil, "", newAssetError(AssetErrorServer, err.Error(), http.StatusInternalServerError)
}
return baseURL + "/api/assets/upload", data, http.MethodPost, nil
case AssetOperationAssetGet:
id := tianyiYunAssetField(req.Body, "Id", "id")
if id == "" {
return "", nil, "", newAssetError(AssetErrorInvalidRequest, "asset id is required", http.StatusBadRequest)
}
return baseURL + "/api/assets/" + url.PathEscape(id), nil, http.MethodGet, nil
default:
return "", nil, "", newAssetError(AssetErrorOperationNotSupported, fmt.Sprintf("asset operation %s is not supported", req.Action.Operation), http.StatusBadRequest)
}
}

func tianyiYunAssetBaseURL(channel *model.Channel) (string, error) {
baseURL := strings.TrimSpace(channel.GetBaseURL())
if baseURL == "" {
baseURL = "https://ai.ctaigw.cn/v1"
}
u, err := url.Parse(baseURL)
if err != nil || u.Scheme == "" || u.Host == "" {
return "", fmt.Errorf("invalid TianyiYun asset base URL")
}
u.Path = strings.TrimRight(u.Path, "/")
if !strings.HasSuffix(u.Path, "/v1") {
u.Path += "/v1"
}
u.RawQuery = ""
u.Fragment = ""
return strings.TrimRight(u.String(), "/"), nil
}

func buildTianyiYunAssetCreatePayload(body map[string]any) (map[string]string, *AssetError) {
sourceURL := tianyiYunAssetField(body, "URL", "url")
if sourceURL == "" {
return nil, newAssetError(AssetErrorInvalidRequest, "asset URL is required", http.StatusBadRequest)
}
parsedURL, err := url.Parse(sourceURL)
if err != nil || (parsedURL.Scheme != "http" && parsedURL.Scheme != "https") || parsedURL.Host == "" {
return nil, newAssetError(AssetErrorInvalidRequest, "asset URL must use http or https", http.StatusBadRequest)
}
assetType := tianyiYunAssetField(body, "AssetType", "asset_type")
if assetType != "Image" && assetType != "Video" && assetType != "Audio" {
return nil, newAssetError(AssetErrorInvalidRequest, "asset type must be Image, Video, or Audio", http.StatusBadRequest)
}
fetchSetting := system_setting.GetFetchSetting()
if err := common.ValidateURLWithFetchSetting(sourceURL, fetchSetting.EnableSSRFProtection, fetchSetting.AllowPrivateIp, fetchSetting.DomainFilterMode, fetchSetting.IpFilterMode, fetchSetting.DomainList, fetchSetting.IpList, fetchSetting.AllowedPorts, fetchSetting.ApplyIPFilterForDomain); err != nil {
return nil, newAssetError(AssetErrorInvalidRequest, fmt.Sprintf("asset URL blocked: %v", err), http.StatusBadRequest)
}
payload := map[string]string{"url": sourceURL, "asset_type": assetType}
if name := tianyiYunAssetField(body, "Name", "name"); name != "" {
payload["name"] = name
}
return payload, nil
}

func tianyiYunAssetField(body map[string]any, upperKey string, lowerKey string) string {
if body == nil {
return ""
}
if value, ok := body[upperKey]; ok {
if text, ok := value.(string); ok {
return strings.TrimSpace(text)
}
}
if value, ok := body[lowerKey]; ok {
if text, ok := value.(string); ok {
return strings.TrimSpace(text)
}
}
return ""
}

+ 131
- 0
service/asset_tianyiyun_test.go Просмотреть файл

@@ -0,0 +1,131 @@
package service

import (
"context"
"io"
"net/http"
"net/http/httptest"
"strings"
"testing"

"github.com/QuantumNous/new-api/common"
"github.com/QuantumNous/new-api/model"
"github.com/QuantumNous/new-api/setting/system_setting"
"github.com/stretchr/testify/require"
)

func TestTianyiYunAssetAdapterCreateAssetMapsActionRequest(t *testing.T) {
disableSSRFProtectionForTianyiYunAssetTest(t)
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
require.Equal(t, http.MethodPost, r.Method)
require.Equal(t, "/v1/api/assets/upload", r.URL.Path)
require.Equal(t, "Bearer key", r.Header.Get("Authorization"))
body, err := io.ReadAll(r.Body)
require.NoError(t, err)
require.JSONEq(t, `{"url":"https://example.test/a.png","asset_type":"Image","name":"cover"}`, string(body))
_, _ = w.Write([]byte(`{"code":0,"data":{"Id":"asset-1"}}`))
}))
defer server.Close()

adapter := NewTianyiYunAssetAdapter()
channel := &model.Channel{BaseURL: common.GetPointer(server.URL + "/v1"), Key: "key"}
response, assetErr := adapter.DoAssetRequest(context.Background(), channel, AssetRequest{
Action: AssetActionSpec{Action: "CreateAsset", Operation: AssetOperationAssetCreate},
Body: map[string]any{"URL": "https://example.test/a.png", "AssetType": "Image", "Name": "cover"},
})

require.Nil(t, assetErr)
require.Equal(t, http.StatusOK, response.StatusCode)
require.JSONEq(t, `{"code":0,"data":{"Id":"asset-1"}}`, string(response.Body))
}

func TestTianyiYunAssetAdapterGetAssetMapsActionRequest(t *testing.T) {
disableSSRFProtectionForTianyiYunAssetTest(t)
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
require.Equal(t, http.MethodGet, r.Method)
require.Equal(t, "/v1/api/assets/asset-1%2Fpart", r.URL.EscapedPath())
require.Equal(t, "Bearer key", r.Header.Get("Authorization"))
body, err := io.ReadAll(r.Body)
require.NoError(t, err)
require.Empty(t, body)
_, _ = w.Write([]byte(`{"code":0,"data":{"Id":"asset-1"}}`))
}))
defer server.Close()

adapter := NewTianyiYunAssetAdapter()
channel := &model.Channel{BaseURL: common.GetPointer(server.URL), Key: "key"}
response, assetErr := adapter.DoAssetRequest(context.Background(), channel, AssetRequest{
Action: AssetActionSpec{Action: "GetAsset", Operation: AssetOperationAssetGet},
Body: map[string]any{"Id": "asset-1/part"},
})

require.Nil(t, assetErr)
require.Equal(t, http.StatusOK, response.StatusCode)
}

func TestTianyiYunAssetAdapterRejectsInvalidCreateRequest(t *testing.T) {
cases := []map[string]any{
{"AssetType": "Image"},
{"URL": "data:image/png;base64,AAAA", "AssetType": "Image"},
{"URL": "aGVsbG8=", "AssetType": "Image"},
{"URL": "ftp://example.test/a.png", "AssetType": "Image"},
{"URL": "https://example.test/a.png", "AssetType": "Document"},
}
for _, body := range cases {
t.Run("invalid", func(t *testing.T) {
adapter := NewTianyiYunAssetAdapter()
response, assetErr := adapter.DoAssetRequest(context.Background(), &model.Channel{Key: "key"}, AssetRequest{
Action: AssetActionSpec{Action: "CreateAsset", Operation: AssetOperationAssetCreate}, Body: body,
})
require.Nil(t, response)
require.NotNil(t, assetErr)
require.Equal(t, http.StatusBadRequest, assetErr.HTTPStatus)
})
}
}

func disableSSRFProtectionForTianyiYunAssetTest(t *testing.T) {
t.Helper()
setting := system_setting.GetFetchSetting()
old := *setting
setting.EnableSSRFProtection = false
t.Cleanup(func() { *setting = old })
}

func TestTianyiYunAssetAdapterRejectsPrivateSourceURL(t *testing.T) {
setting := system_setting.GetFetchSetting()
old := *setting
setting.EnableSSRFProtection = true
setting.AllowPrivateIp = false
setting.DomainFilterMode = false
setting.IpFilterMode = false
setting.DomainList = nil
setting.IpList = nil
setting.AllowedPorts = []string{"80", "443", "8080", "8443"}
t.Cleanup(func() { *setting = old })

for _, sourceURL := range []string{"https://127.0.0.1/a.png", "http://10.0.0.1/a.png"} {
t.Run(sourceURL, func(t *testing.T) {
response, assetErr := NewTianyiYunAssetAdapter().DoAssetRequest(context.Background(), &model.Channel{Key: "key"}, AssetRequest{
Action: AssetActionSpec{Action: "CreateAsset", Operation: AssetOperationAssetCreate},
Body: map[string]any{"URL": sourceURL, "AssetType": "Image"},
})
require.Nil(t, response)
require.NotNil(t, assetErr)
require.Equal(t, http.StatusBadRequest, assetErr.HTTPStatus)
})
}
}

func TestTianyiYunAssetAdapterSupportsOnlyCreateAndGet(t *testing.T) {
adapter := NewTianyiYunAssetAdapter()
require.True(t, adapter.Supports(AssetOperationAssetCreate))
require.True(t, adapter.Supports(AssetOperationAssetGet))
require.False(t, adapter.Supports(AssetOperationAssetList))
_, assetErr := adapter.DoAssetRequest(context.Background(), &model.Channel{Key: "key"}, AssetRequest{
Action: AssetActionSpec{Operation: AssetOperationAssetList},
})
require.NotNil(t, assetErr)
require.Equal(t, AssetErrorOperationNotSupported, assetErr.Type)
require.True(t, strings.Contains(assetErr.Message, "not supported"))
}

+ 51
- 0
service/audio_test.go Просмотреть файл

@@ -0,0 +1,51 @@
package service

import (
"encoding/base64"
"testing"

"github.com/stretchr/testify/require"
)

func TestParseAudioUsesFormatSpecificSampleRates(t *testing.T) {
pcm := base64.StdEncoding.EncodeToString(make([]byte, 48_000))
duration, err := parseAudio(pcm, "pcm16")
require.NoError(t, err)
require.Equal(t, 1.0, duration)

g711 := base64.StdEncoding.EncodeToString(make([]byte, 8_000))
duration, err = parseAudio(g711, "g711_ulaw")
require.NoError(t, err)
require.Equal(t, 1.0, duration)
}

func TestParseAudioAndDecodeBase64AudioDataRejectInvalidEncoding(t *testing.T) {
_, err := parseAudio("not base64", "pcm16")
require.Error(t, err)

_, err = DecodeBase64AudioData("not base64")
require.Error(t, err)
}

func TestDecodeBase64AudioDataStripsDataURLPrefix(t *testing.T) {
decoded, err := DecodeBase64AudioData("data:audio/pcm;base64,AAE=")

require.NoError(t, err)
require.Equal(t, "AAE=", decoded)
}

func TestCountAudioTokensUsesConfiguredInputAndOutputRates(t *testing.T) {
oneSecondPCM := base64.StdEncoding.EncodeToString(make([]byte, 48_000))

inputTokens, err := CountAudioTokenInput(oneSecondPCM, "pcm16")
require.NoError(t, err)
require.Equal(t, 27, inputTokens)

outputTokens, err := CountAudioTokenOutput(oneSecondPCM, "pcm16")
require.NoError(t, err)
require.Equal(t, 13, outputTokens)

inputTokens, err = CountAudioTokenInput("", "pcm16")
require.NoError(t, err)
require.Zero(t, inputTokens)
}

+ 35
- 0
service/channel_affinity_helpers_test.go Просмотреть файл

@@ -0,0 +1,35 @@
package service

import (
"testing"

"github.com/QuantumNous/new-api/dto"
"github.com/QuantumNous/new-api/setting/operation_setting"
"github.com/QuantumNous/new-api/types"
"github.com/stretchr/testify/require"
)

func TestChannelAffinityMatchersIgnoreInvalidPatternsAndMatchCaseInsensitively(t *testing.T) {
require.True(t, matchAnyRegexCached([]string{"[", `^gpt-[0-9]+$`}, "gpt-5"))
require.False(t, matchAnyRegexCached([]string{"["}, "gpt-5"))
require.True(t, matchAnyIncludeFold([]string{" Codex ", "other"}, "Mozilla CodexClient"))
require.False(t, matchAnyIncludeFold([]string{""}, "CodexClient"))
}

func TestChannelAffinityKeyHelpersProduceStableSafeValues(t *testing.T) {
rule := operation_setting.ChannelAffinityRule{Name: "rule", IncludeRuleName: true, IncludeUsingGroup: true}
require.Equal(t, "rule:vip:key", buildChannelAffinityCacheKeySuffix(rule, "vip", "key"))
require.Equal(t, "abcd...wxyz", buildChannelAffinityKeyHint("abcdefghijklmnopqrstuvwxwxyz"))
require.Len(t, affinityFingerprint("tenant-123"), 8)
require.Equal(t, "", channelAffinityUsageCacheEntryKey("", "vip", "fingerprint"))
require.Equal(t, "rule\n\nfingerprint", channelAffinityUsageCacheEntryKey("rule", "", "fingerprint"))
}

func TestChannelAffinityUsageHelpersUseFallbackFields(t *testing.T) {
usage := &dto.Usage{InputTokens: 3, OutputTokens: 5}
require.Equal(t, 3, usagePromptTokens(usage))
require.Equal(t, 5, usageCompletionTokens(usage))
require.Equal(t, 8, usageTotalTokens(usage))
require.Equal(t, cacheTokenRateModeCachedOverPrompt, cachedTokenRateModeByRelayFormat(types.RelayFormatOpenAI))
require.Equal(t, cacheTokenRateModeCachedOverPromptPlusCached, cachedTokenRateModeByRelayFormat(types.RelayFormatClaude))
}

+ 20
- 0
service/channel_select_test.go Просмотреть файл

@@ -179,6 +179,26 @@ func TestCacheGetRandomSatisfiedChannelFiltersAllowedChannelTypes(t *testing.T)
assert.Equal(t, 101, channel.Id)
}

func TestCacheGetRandomSatisfiedChannelAllowsChinaMobileSeedanceFamily(t *testing.T) {
db := setupServiceChannelSelectDB(t)
createServiceChannelSelectChannel(t, db, 101, constant.ChannelTypeChinaMobileSeedance, "default", "doubao-seedance-2.0")
createServiceChannelSelectChannel(t, db, 102, constant.ChannelTypeOpenAI, "default", "doubao-seedance-2.0")

c := buildRetryContext(t)
channel, group, err := CacheGetRandomSatisfiedChannel(&RetryParam{
Ctx: c,
TokenGroup: "default",
ModelName: "doubao-seedance-2.0",
Retry: common.GetPointer(0),
AllowedChannelTypes: VideoAssetChannelTypesForFamily(VideoAssetFamilySeedance),
})

require.NoError(t, err)
require.NotNil(t, channel)
assert.Equal(t, "default", group)
assert.Equal(t, 101, channel.Id)
}

// --- GetUserAutoGroup tests ---

func TestGetUserAutoGroup_Basic(t *testing.T) {


+ 67
- 0
service/channel_test.go Просмотреть файл

@@ -0,0 +1,67 @@
package service

import (
"errors"
"net/http"
"testing"

"github.com/QuantumNous/new-api/common"
"github.com/QuantumNous/new-api/constant"
"github.com/QuantumNous/new-api/setting/operation_setting"
"github.com/QuantumNous/new-api/types"
"github.com/stretchr/testify/require"
)

func TestShouldDisableChannelHonorsFeatureFlagAndChannelErrors(t *testing.T) {
original := common.AutomaticDisableChannelEnabled
t.Cleanup(func() { common.AutomaticDisableChannelEnabled = original })

channelErr := types.NewError(errors.New("no usable key"), types.ErrorCodeChannelNoAvailableKey)
common.AutomaticDisableChannelEnabled = false
require.False(t, ShouldDisableChannel(constant.ChannelTypeOpenAI, channelErr))

common.AutomaticDisableChannelEnabled = true
require.False(t, ShouldDisableChannel(constant.ChannelTypeOpenAI, nil))
require.True(t, ShouldDisableChannel(constant.ChannelTypeOpenAI, channelErr))

skipRetry := types.NewError(errors.New("retry later"), types.ErrorCodeBadResponse, types.ErrOptionWithSkipRetry())
require.False(t, ShouldDisableChannel(constant.ChannelTypeOpenAI, skipRetry))
}

func TestShouldDisableChannelRecognizesStatusAndOpenAIErrorRules(t *testing.T) {
originalEnabled := common.AutomaticDisableChannelEnabled
originalRanges := operation_setting.AutomaticDisableStatusCodeRanges
t.Cleanup(func() {
common.AutomaticDisableChannelEnabled = originalEnabled
operation_setting.AutomaticDisableStatusCodeRanges = originalRanges
})
common.AutomaticDisableChannelEnabled = true
operation_setting.AutomaticDisableStatusCodeRanges = []operation_setting.StatusCodeRange{{Start: http.StatusTooManyRequests, End: http.StatusTooManyRequests}}

byStatus := types.NewOpenAIError(errors.New("rate limited"), types.ErrorCodeBadResponse, http.StatusTooManyRequests)
require.True(t, ShouldDisableChannel(constant.ChannelTypeOpenAI, byStatus))

forbidden := types.NewOpenAIError(errors.New("forbidden"), types.ErrorCodeBadResponse, http.StatusForbidden)
require.True(t, ShouldDisableChannel(constant.ChannelTypeGemini, forbidden))
require.False(t, ShouldDisableChannel(constant.ChannelTypeOpenAI, forbidden))

invalidKey := types.WithOpenAIError(types.OpenAIError{
Message: "bad key",
Type: "invalid_request_error",
Code: "invalid_api_key",
}, http.StatusBadRequest)
require.True(t, ShouldDisableChannel(constant.ChannelTypeOpenAI, invalidKey))
}

func TestShouldEnableChannelRequiresAutoDisabledStatusWithoutError(t *testing.T) {
original := common.AutomaticEnableChannelEnabled
t.Cleanup(func() { common.AutomaticEnableChannelEnabled = original })

common.AutomaticEnableChannelEnabled = false
require.False(t, ShouldEnableChannel(nil, common.ChannelStatusAutoDisabled))

common.AutomaticEnableChannelEnabled = true
require.False(t, ShouldEnableChannel(types.NewError(errors.New("still failing"), types.ErrorCodeBadResponse), common.ChannelStatusAutoDisabled))
require.False(t, ShouldEnableChannel(nil, common.ChannelStatusManuallyDisabled))
require.True(t, ShouldEnableChannel(nil, common.ChannelStatusAutoDisabled))
}

+ 128
- 0
service/chinamobile_user_asset_group.go Просмотреть файл

@@ -0,0 +1,128 @@
package service

import (
"context"
"fmt"
"net/http"
"strings"

"github.com/QuantumNous/new-api/common"
"github.com/QuantumNous/new-api/constant"
"github.com/QuantumNous/new-api/model"
)

type chinaMobileAssetGroupCreator func(context.Context, int, *model.Channel) (string, *AssetError)

func IsChinaMobileAssetChannel(channel *model.Channel) bool {
return channel != nil && channel.Type == constant.ChannelTypeChinaMobileSeedance
}

func IsAssetGroupOperation(operation AssetOperation) bool {
return strings.HasPrefix(string(operation), "asset_group.")
}

func GetOrCreateChinaMobileUserAssetGroup(ctx context.Context, userId int, channel *model.Channel, createGroup chinaMobileAssetGroupCreator) (string, *AssetError) {
if channel == nil {
return "", newAssetError(AssetErrorServer, "China Mobile asset channel is required", http.StatusInternalServerError)
}
binding, err := model.GetUserAssetGroup(userId, channel.Id)
if err != nil {
return "", newAssetError(AssetErrorServer, err.Error(), http.StatusInternalServerError)
}
if binding != nil {
return binding.GroupId, nil
}

groupId, assetErr := createGroup(ctx, userId, channel)
if assetErr != nil {
return "", assetErr
}
if strings.TrimSpace(groupId) == "" {
return "", newAssetError(AssetErrorUpstream, "China Mobile asset group creation returned an empty group ID", http.StatusBadGateway)
}
if err := model.CreateUserAssetGroup(userId, channel.Id, groupId); err == nil {
return groupId, nil
}

// Another first request may have persisted the binding while this one created its upstream group.
binding, readErr := model.GetUserAssetGroup(userId, channel.Id)
if readErr == nil && binding != nil {
return binding.GroupId, nil
}
if readErr != nil {
return "", newAssetError(AssetErrorServer, readErr.Error(), http.StatusInternalServerError)
}
return "", newAssetError(AssetErrorServer, "failed to persist China Mobile user asset group", http.StatusInternalServerError)
}

func CreateChinaMobileUserAssetGroup(ctx context.Context, userId int, adapter AssetAdapter, channel *model.Channel) (string, *AssetError) {
spec, ok := ParseAssetAction("CreateAssetGroup")
if !ok {
return "", newAssetError(AssetErrorServer, "CreateAssetGroup action is not registered", http.StatusInternalServerError)
}
resp, assetErr := adapter.DoAssetRequest(ctx, channel, AssetRequest{
Action: spec,
Version: "2024-01-01",
Body: map[string]any{
"GroupType": "AIGC",
"Name": fmt.Sprintf("new-api-user-%d-channel-%d", userId, channel.Id),
},
})
if assetErr != nil {
return "", assetErr
}
var payload struct {
Result struct {
GroupId string `json:"GroupId"`
} `json:"Result"`
}
if err := common.Unmarshal(resp.Body, &payload); err != nil {
return "", newAssetError(AssetErrorUpstream, fmt.Sprintf("invalid China Mobile asset group response: %v", err), http.StatusBadGateway)
}
return strings.TrimSpace(payload.Result.GroupId), nil
}

func ScopeChinaMobileAssetRequest(req *AssetRequest, groupId string) {
switch req.Action.Operation {
case AssetOperationAssetCreate:
req.Body["GroupId"] = groupId
case AssetOperationAssetList:
filter := mapValue(req.Body, "Filter")
if filter == nil {
return
}
filter["GroupIds"] = []string{groupId}
req.Body["Filter"] = filter
}
}

func RequireChinaMobileAssetOwnership(ctx context.Context, adapter AssetAdapter, channel *model.Channel, request AssetRequest, groupId string) *AssetError {
spec, ok := ParseAssetAction("GetAsset")
if !ok {
return newAssetError(AssetErrorServer, "GetAsset action is not registered", http.StatusInternalServerError)
}
assetId := strings.TrimSpace(stringValue(request.Body, "Id"))
resp, assetErr := adapter.DoAssetRequest(ctx, channel, AssetRequest{
Action: spec,
Version: request.Version,
Body: map[string]any{"Id": assetId},
})
if assetErr != nil {
if assetErr.HTTPStatus == http.StatusNotFound {
return newAssetError(AssetErrorNotFound, "asset not found", http.StatusNotFound)
}
return assetErr
}
var payload struct {
Result struct {
GroupId string `json:"GroupId"`
} `json:"Result"`
}
if err := common.Unmarshal(resp.Body, &payload); err != nil {
return newAssetError(AssetErrorUpstream, fmt.Sprintf("invalid China Mobile asset response: %v", err), http.StatusBadGateway)
}
if strings.TrimSpace(payload.Result.GroupId) != groupId {
return newAssetError(AssetErrorNotFound, "asset not found", http.StatusNotFound)
}
return nil
}

+ 159
- 0
service/chinamobile_user_asset_group_test.go Просмотреть файл

@@ -0,0 +1,159 @@
package service

import (
"context"
"net/http"
"sync"
"testing"

"github.com/QuantumNous/new-api/common"
"github.com/QuantumNous/new-api/constant"
"github.com/QuantumNous/new-api/model"
"github.com/glebarez/sqlite"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"gorm.io/gorm"
"gorm.io/gorm/logger"
)

func setupChinaMobileUserAssetGroupDB(t *testing.T) *gorm.DB {
t.Helper()
db, err := gorm.Open(sqlite.Open("file:chinamobile_user_asset_groups?mode=memory&cache=shared"), &gorm.Config{Logger: logger.Default.LogMode(logger.Silent)})
require.NoError(t, err)
sqlDB, err := db.DB()
require.NoError(t, err)
sqlDB.SetMaxOpenConns(1)

origDB := model.DB
model.DB = db
require.NoError(t, db.AutoMigrate(&model.UserAssetGroup{}))
t.Cleanup(func() {
model.DB = origDB
require.NoError(t, sqlDB.Close())
})
return db
}

func chinaMobileAssetAction(t *testing.T, action string) AssetActionSpec {
t.Helper()
spec, ok := ParseAssetAction(action)
require.True(t, ok)
return spec
}

func TestGetOrCreateChinaMobileUserAssetGroupCreatesThenReuses(t *testing.T) {
setupChinaMobileUserAssetGroupDB(t)
channel := &model.Channel{Id: 101, Type: constant.ChannelTypeChinaMobileSeedance}
createCalls := 0
creator := func(context.Context, int, *model.Channel) (string, *AssetError) {
createCalls++
return "group-1", nil
}

groupID, assetErr := GetOrCreateChinaMobileUserAssetGroup(context.Background(), 10, channel, creator)
require.Nil(t, assetErr)
assert.Equal(t, "group-1", groupID)

again, assetErr := GetOrCreateChinaMobileUserAssetGroup(context.Background(), 10, channel, creator)
require.Nil(t, assetErr)
assert.Equal(t, "group-1", again)
assert.Equal(t, 1, createCalls)
}

func TestGetOrCreateChinaMobileUserAssetGroupConcurrentRequestsReuseBinding(t *testing.T) {
setupChinaMobileUserAssetGroupDB(t)
channel := &model.Channel{Id: 102, Type: constant.ChannelTypeChinaMobileSeedance}
creator := func(context.Context, int, *model.Channel) (string, *AssetError) {
return "group-concurrent", nil
}

const workers = 8
results := make(chan string, workers)
errs := make(chan *AssetError, workers)
var wg sync.WaitGroup
for i := 0; i < workers; i++ {
wg.Add(1)
go func() {
defer wg.Done()
groupID, assetErr := GetOrCreateChinaMobileUserAssetGroup(context.Background(), 10, channel, creator)
results <- groupID
errs <- assetErr
}()
}
wg.Wait()
close(results)
close(errs)
for assetErr := range errs {
require.Nil(t, assetErr)
}
for groupID := range results {
assert.Equal(t, "group-concurrent", groupID)
}

binding, err := model.GetUserAssetGroup(10, 102)
require.NoError(t, err)
require.NotNil(t, binding)
assert.Equal(t, "group-concurrent", binding.GroupId)
}

func TestScopeChinaMobileAssetRequestOverwritesClientGroup(t *testing.T) {
create := AssetRequest{Action: chinaMobileAssetAction(t, "CreateAsset"), Body: map[string]any{"GroupId": "forged"}}
ScopeChinaMobileAssetRequest(&create, "owned")
assert.Equal(t, "owned", create.Body["GroupId"])

list := AssetRequest{Action: chinaMobileAssetAction(t, "ListAssets"), Body: map[string]any{"Filter": map[string]any{"GroupIds": []any{"forged"}}}}
ScopeChinaMobileAssetRequest(&list, "owned")
assert.Equal(t, []string{"owned"}, list.Body["Filter"].(map[string]any)["GroupIds"])
}

func TestRequireChinaMobileAssetOwnershipHidesMismatchedGroup(t *testing.T) {
adapter := &recordingAssetAdapter{getAssetGroupID: "other"}
request := AssetRequest{Action: chinaMobileAssetAction(t, "DeleteAsset"), Body: map[string]any{"Id": "asset-1"}}

assetErr := RequireChinaMobileAssetOwnership(context.Background(), adapter, &model.Channel{}, request, "owned")

require.NotNil(t, assetErr)
assert.Equal(t, AssetErrorNotFound, assetErr.Type)
assert.Equal(t, http.StatusNotFound, assetErr.HTTPStatus)
assert.Equal(t, AssetOperationAssetGet, adapter.requests[0].Action.Operation)
}

type recordingAssetAdapter struct {
requests []AssetRequest
getAssetGroupID string
}

func (a *recordingAssetAdapter) Name() string { return "recording" }
func (a *recordingAssetAdapter) Supports(AssetOperation) bool { return true }
func (a *recordingAssetAdapter) DoAssetRequest(_ context.Context, _ *model.Channel, req AssetRequest) (*AssetUpstreamResponse, *AssetError) {
a.requests = append(a.requests, req)
if req.Action.Operation == AssetOperationAssetGroupCreate {
body, err := BuildAssetSuccessResponse(req.Action.Action, req.Version, map[string]any{"GroupId": "group-created"})
if err != nil {
return nil, newAssetError(AssetErrorServer, err.Error(), http.StatusInternalServerError)
}
return &AssetUpstreamResponse{StatusCode: http.StatusOK, Body: body}, nil
}
if req.Action.Operation != AssetOperationAssetGet {
return &AssetUpstreamResponse{StatusCode: http.StatusOK}, nil
}
body, err := BuildAssetSuccessResponse(req.Action.Action, req.Version, map[string]any{"Id": "asset-1", "GroupId": a.getAssetGroupID})
if err != nil {
return nil, newAssetError(AssetErrorServer, err.Error(), http.StatusInternalServerError)
}
return &AssetUpstreamResponse{StatusCode: http.StatusOK, Body: body}, nil
}

func TestCreateChinaMobileUserAssetGroupUsesInternalGroupAction(t *testing.T) {
adapter := &recordingAssetAdapter{}
groupID, assetErr := CreateChinaMobileUserAssetGroup(context.Background(), 10, adapter, &model.Channel{Id: 101})
require.Nil(t, assetErr)
assert.Equal(t, "group-created", groupID)
assert.Equal(t, AssetOperationAssetGroupCreate, adapter.requests[0].Action.Operation)
assert.Equal(t, "AIGC", adapter.requests[0].Body["GroupType"])
assert.Equal(t, "new-api-user-10-channel-101", adapter.requests[0].Body["Name"])

// Ensure the normalized response parser remains wired through common JSON helpers.
_, err := common.Marshal(adapter.requests[0].Body)
require.NoError(t, err)
}

+ 6
- 6
service/codex_oauth.go Просмотреть файл

@@ -120,12 +120,12 @@ func refreshCodexOAuthToken(
ExpiresIn int `json:"expires_in"`
}

if err := common.DecodeJson(resp.Body, &payload); err != nil {
return nil, err
}
if resp.StatusCode < 200 || resp.StatusCode >= 300 {
return nil, fmt.Errorf("codex oauth refresh failed: status=%d", resp.StatusCode)
}
if err := common.DecodeJson(resp.Body, &payload); err != nil {
return nil, err
}

if strings.TrimSpace(payload.AccessToken) == "" || strings.TrimSpace(payload.RefreshToken) == "" || payload.ExpiresIn <= 0 {
return nil, errors.New("codex oauth refresh response missing fields")
@@ -181,12 +181,12 @@ func exchangeCodexAuthorizationCode(
RefreshToken string `json:"refresh_token"`
ExpiresIn int `json:"expires_in"`
}
if err := common.DecodeJson(resp.Body, &payload); err != nil {
return nil, err
}
if resp.StatusCode < 200 || resp.StatusCode >= 300 {
return nil, fmt.Errorf("codex oauth code exchange failed: status=%d", resp.StatusCode)
}
if err := common.DecodeJson(resp.Body, &payload); err != nil {
return nil, err
}
if strings.TrimSpace(payload.AccessToken) == "" || strings.TrimSpace(payload.RefreshToken) == "" || payload.ExpiresIn <= 0 {
return nil, errors.New("codex oauth token response missing fields")
}


Некоторые файлы не были показаны из-за большого количества измененных файлов

Загрузка…
Отмена
Сохранить