| Автор | SHA1 | Сообщение | Дата |
|---|---|---|---|
|
|
83724d4930 |
chore: add image build script, drop binary artifact, add missing go.sum
Add deploy/build_newapi.sh for tagged Aliyun image builds, remove the 94MB compiled binary new-api.exe~ accidentally committed to the repo, and add the missing go.sum for third_party/ecloudsdkcore. Co-Authored-By: ZCode <noreply@anthropic.com> |
1 месяц назад |
|
|
ce40bfe92f |
feat(chinamobile): support doubao-seedance-2-0-mini-260615 billing model
Register the Seedance mini model in the matrix usage capability table (both native name and doubao-seedance-2.0 alias) and keep SupportsAnyMatrixUsageBillingModel in sync so the model bills via matrix usage on the ChinaMobile Seedance channel. Co-Authored-By: ZCode <noreply@anthropic.com> |
1 месяц назад |
|
|
0b79e4ee07 |
fix(db): skip logs AutoMigrate for pg_partman-managed PostgreSQL tables
On PostgreSQL the logs table is a RANGE-partitioned table on created_at owned by pg_partman; GORM AutoMigrate would alter its composite primary key (id, created_at) and create redundant per-partition indexes, breaking startup. Extract ensureLogTable() to skip AutoMigrate on PostgreSQL (existence check only) while keeping the original behavior on SQLite/MySQL. The dedicated LOG_SQL_DSN PostgreSQL branch follows the same rule. Ship the matching postgres-partman image (Dockerfile + 01-partition.sql) that creates and maintains the partitioned table via pg_partman + pg_cron. Co-Authored-By: ZCode <noreply@anthropic.com> |
1 месяц назад |
|
|
17c12fd41e |
feat: harden services and add video channel bindings
Co-Authored-By: Codex <noreply@anthropic.com> |
2 месяцев назад |
|
|
cd35e66b33 |
feat(chinamobile): add channel asset credentials
Co-Authored-By: Codex <noreply@anthropic.com> |
2 месяцев назад |
|
|
50ce4f7541 |
feat(tianyiyun): add Seedance asset support
Co-Authored-By: Codex <noreply@anthropic.com> |
2 месяцев назад |
|
|
fedd394e79 |
fix(task-billing): persist settlement and merge paged logs
Co-Authored-By: Codex <noreply@anthropic.com> |
2 месяцев назад |
|
|
c67d8c409e |
feat(chinamobile): add Seedance channel and isolated asset library
Co-Authored-By: Codex <noreply@anthropic.com> |
2 месяцев назад |
|
|
fcb1c2a8ba |
fix: prevent webhook bypass via empty secret and cross-gateway callback attacks
Add two layers of defense:
1. Webhook availability guard (controller/payment_webhook_availability.go):
- isStripeWebhookEnabled() checks StripeWebhookSecret != "" before processing
- isCreemWebhookEnabled(), isWechatPayWebhookEnabled(), isAlipayWebhookEnabled()
- isEpayWebhookEnabled() with similar checks for all payment webhooks
- Applied to: StripeWebhook, CreemWebhook, WechatPayWebhook, AlipayPayWebhook,
EpayNotify, SubscriptionEpayNotify
2. PaymentProvider field (model/topup.go):
- New PaymentProvider field on TopUp to identify which gateway created the order
- Recharge() checks PaymentProvider == PaymentProviderStripe
- rechargeByQRCodePayment() checks PaymentProvider matches wechat/alipay
- RechargeCreem() checks PaymentProvider == PaymentProviderCreem
- All payment controllers set PaymentProvider when creating orders
Root cause: When StripeWebhookSecret was empty, ComputeSignature used
an empty HMAC key, allowing attackers to forge valid signatures and
complete orders from any payment gateway without actually paying.
Co-Authored-By: Claude <noreply@anthropic.com>
|
2 месяцев назад |
|
|
6891b11e09 |
chore: ignore local tmp artifacts
Co-Authored-By: Codex <noreply@anthropic.com> |
2 месяцев назад |
|
|
bde79b7259 |
fix(kling): persist proxy fallback asset binding
Co-Authored-By: Codex <noreply@anthropic.com> |
2 месяцев назад |
| @@ -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/ | |||
| @@ -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 . . | |||
| @@ -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: | |||
| @@ -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) | |||
| } | |||
| @@ -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) { | |||
| @@ -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 { | |||
| @@ -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) { | |||
| @@ -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) { | |||
| }, | |||
| }) | |||
| } | |||
| @@ -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`) | |||
| } | |||
| @@ -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") | |||
| } | |||
| @@ -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) | |||
| } | |||
| @@ -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) | |||
| } | |||
| @@ -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())) | |||
| } | |||
| } | |||
| @@ -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 | |||
| } | |||
| @@ -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 | |||
| } | |||
| @@ -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")) | |||
| } | |||
| @@ -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 { | |||
| @@ -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) | |||
| } | |||
| @@ -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") | |||
| } | |||
| @@ -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 | |||
| } | |||
| @@ -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()) | |||
| } | |||
| @@ -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()) | |||
| } | |||
| @@ -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" { | |||
| @@ -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" { | |||
| @@ -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) | |||
| @@ -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 { | |||
| @@ -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) | |||
| @@ -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) | |||
| @@ -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, | |||
| @@ -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": ""}) | |||
| } | |||
| @@ -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) | |||
| } | |||
| } | |||
| @@ -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", "")) | |||
| } | |||
| @@ -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" | |||
| @@ -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 | |||
| @@ -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()$$ | |||
| ); | |||
| @@ -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. | |||
| @@ -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` 持久化全部验证通过。 | |||
| @@ -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` 未导出该符号。新增中英文键已手动同步,生产构建已验证通过。 | |||
| @@ -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 | |||
| @@ -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= | |||
| @@ -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) | |||
| } | |||
| @@ -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)) | |||
| } | |||
| @@ -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) | |||
| } | |||
| @@ -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) { | |||
| @@ -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 | |||
| } | |||
| @@ -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) | |||
| } | |||
| @@ -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") | |||
| } | |||
| @@ -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 | |||
| } | |||
| @@ -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) | |||
| @@ -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 | |||
| @@ -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()) | |||
| } | |||
| @@ -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) | |||
| } | |||
| @@ -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) | |||
| } | |||
| @@ -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")) | |||
| } | |||
| @@ -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) | |||
| @@ -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()) | |||
| } | |||
| @@ -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 | |||
| @@ -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("无效的充值额度") | |||
| @@ -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()) | |||
| } | |||
| @@ -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 | |||
| @@ -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) | |||
| @@ -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 | |||
| } | |||
| @@ -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) | |||
| } | |||
| @@ -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)) | |||
| } | |||
| @@ -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)) | |||
| } | |||
| @@ -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" | |||
| } | |||
| @@ -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 | |||
| } | |||
| @@ -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" | |||
| @@ -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) | |||
| } | |||
| } | |||
| @@ -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 { | |||
| @@ -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, `{ | |||
| @@ -3,6 +3,7 @@ package doubao_tianyiyun | |||
| var ModelList = []string{ | |||
| "cdance2.0-0611", | |||
| "cdance2.0-fast-0611", | |||
| "cdance2.0-mini-0611", | |||
| } | |||
| var ChannelName = "DoubaoVideoCompatibleTianyiYun" | |||
| @@ -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}, | |||
| @@ -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 | |||
| @@ -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: | |||
| @@ -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()) | |||
| } | |||
| @@ -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 | |||
| @@ -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) | |||
| @@ -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 | |||
| } | |||
| @@ -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()) | |||
| } | |||
| @@ -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 | |||
| } | |||
| @@ -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), | |||
| } | |||
| } | |||
| @@ -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) | |||
| } | |||
| @@ -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 | |||
| } | |||
| @@ -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 | |||
| } | |||
| @@ -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) | |||
| } | |||
| @@ -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) | |||
| } | |||
| @@ -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"`) | |||
| } | |||
| @@ -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 | |||
| } | |||
| @@ -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)) | |||
| } | |||
| @@ -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 "" | |||
| } | |||
| @@ -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")) | |||
| } | |||
| @@ -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) | |||
| } | |||
| @@ -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)) | |||
| } | |||
| @@ -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) { | |||
| @@ -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)) | |||
| } | |||
| @@ -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 | |||
| } | |||
| @@ -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) | |||
| } | |||
| @@ -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") | |||
| } | |||