package doubao_tianyiyun import ( "io" "net/http" "net/http/httptest" "strings" "testing" "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 TestBuildRequestURLUsesTianyiYunEndpoint(t *testing.T) { adaptor := &TaskAdaptor{} adaptor.Init(&relaycommon.RelayInfo{ ChannelMeta: &relaycommon.ChannelMeta{ ChannelType: constant.ChannelTypeDoubaoVideoCompatibleTianyiYun, ChannelBaseUrl: "https://ai.ctaigw.cn", ApiKey: "sk-test", }, }) got, err := adaptor.BuildRequestURL(&relaycommon.RelayInfo{ ChannelMeta: &relaycommon.ChannelMeta{ChannelBaseUrl: "https://ai.ctaigw.cn"}, }) require.NoError(t, err) require.Equal(t, "https://ai.ctaigw.cn/v1/contents/generations/tasks", got) } func TestFetchTaskUsesTianyiYunEndpoint(t *testing.T) { adaptor := &TaskAdaptor{} req, err := adaptor.buildFetchRequest("https://ai.ctaigw.cn", "sk-test", map[string]any{"task_id": "task-123"}) require.NoError(t, err) require.Equal(t, http.MethodGet, req.Method) require.Equal(t, "https://ai.ctaigw.cn/v1/contents/generations/tasks/task-123", req.URL.String()) require.Equal(t, "Bearer sk-test", req.Header.Get("Authorization")) require.Equal(t, "application/json", req.Header.Get("Accept")) } func TestGetModelListIncludesTianyiYunSeedanceModels(t *testing.T) { models := (&TaskAdaptor{}).GetModelList() require.Contains(t, models, "cdance2.0-0611") require.Contains(t, models, "cdance2.0-fast-0611") require.Equal(t, "DoubaoVideoCompatibleTianyiYun", (&TaskAdaptor{}).GetChannelName()) } func TestBuildRequestBodyPreservesTianyiYunNativeFields(t *testing.T) { adaptor := &TaskAdaptor{} c := newTianyiYunTaskRequestContext(t, `{ "model":"cdance2.0-0611", "prompt":"write a clean product video", "metadata":{ "content":[ {"type":"image_url","image_url":{"url":"https://example.test/input.png"},"role":"reference_image"} ], "ratio":"16:9", "duration":5, "watermark":false } }`) info := &relaycommon.RelayInfo{ ChannelMeta: &relaycommon.ChannelMeta{UpstreamModelName: "cdance2.0-0611"}, } body, err := adaptor.BuildRequestBody(c, info) require.NoError(t, err) data, err := io.ReadAll(body) require.NoError(t, err) require.Contains(t, string(data), `"model":"cdance2.0-0611"`) require.Contains(t, string(data), `"image_url":{"url":"https://example.test/input.png"}`) require.Contains(t, string(data), `"text":"write a clean product video"`) require.Contains(t, string(data), `"ratio":"16:9"`) require.Contains(t, string(data), `"duration":5`) require.Contains(t, string(data), `"watermark":false`) require.NotContains(t, string(data), `"seconds"`) } func TestBuildRequestBodyUsesMappedTianyiYunUpstreamModel(t *testing.T) { adaptor := &TaskAdaptor{} c := newTianyiYunTaskRequestContext(t, `{ "model":"Doubao-Seedance-2.0", "prompt":"write a clean product video" }`) info := &relaycommon.RelayInfo{ OriginModelName: "Doubao-Seedance-2.0", ChannelMeta: &relaycommon.ChannelMeta{ UpstreamModelName: "cdance2.0-0611", IsModelMapped: true, }, TaskRelayInfo: &relaycommon.TaskRelayInfo{}, } body, err := adaptor.BuildRequestBody(c, info) require.NoError(t, err) data, err := io.ReadAll(body) require.NoError(t, err) require.Contains(t, string(data), `"model":"cdance2.0-0611"`) require.NotContains(t, string(data), `"model":"Doubao-Seedance-2.0"`) } func TestDoResponseRewritesIDToPublicTaskID(t *testing.T) { adaptor := &TaskAdaptor{} w := httptest.NewRecorder() c, _ := gin.CreateTestContext(w) info := &relaycommon.RelayInfo{ OriginModelName: "cdance2.0-0611", TaskRelayInfo: &relaycommon.TaskRelayInfo{PublicTaskID: "task_public"}, } resp := &http.Response{ StatusCode: http.StatusOK, Body: io.NopCloser(strings.NewReader(`{ "id":"task-upstream", "status":"queued", "created_at":1781496040, "model":"cdance2.0-0611" }`)), } upstreamID, data, taskErr := adaptor.DoResponse(c, resp, info) require.Nil(t, taskErr) require.Equal(t, "task-upstream", upstreamID) require.Contains(t, string(data), `"id":"task-upstream"`) require.Equal(t, http.StatusOK, w.Code) require.Contains(t, w.Body.String(), `"id":"task_public"`) require.Contains(t, w.Body.String(), `"status":"queued"`) require.Contains(t, w.Body.String(), `"model":"cdance2.0-0611"`) } func TestParseTaskResultMapsTianyiYunStatuses(t *testing.T) { adaptor := &TaskAdaptor{} cases := []struct { name string body string wantStatus model.TaskStatus wantURL string wantReason string wantUsage bool }{ { name: "success", body: `{"id":"task-ok","status":"success","content":{"video_url":"https://example.test/out.mp4"},"usage":{"completion_tokens":7,"total_tokens":11}}`, wantStatus: model.TaskStatusSuccess, wantURL: "https://example.test/out.mp4", wantUsage: true, }, { name: "succeeded", body: `{"id":"task-ok","status":"succeeded","content":{"video_url":"https://example.test/out.mp4"}}`, wantStatus: model.TaskStatusSuccess, wantURL: "https://example.test/out.mp4", }, { name: "running", body: `{"id":"task-run","status":"running"}`, wantStatus: model.TaskStatusInProgress, }, { name: "failed", body: `{"id":"task-fail","status":"failed","error":{"code":"BadRequest","message":"bad prompt"}}`, wantStatus: model.TaskStatusFailure, wantReason: "bad prompt", }, } for _, tc := range cases { t.Run(tc.name, func(t *testing.T) { got, err := adaptor.ParseTaskResult([]byte(tc.body)) require.NoError(t, err) require.Equal(t, string(tc.wantStatus), got.Status) if tc.wantURL != "" { require.Equal(t, tc.wantURL, got.Url) } if tc.wantUsage { require.Equal(t, 7, got.CompletionTokens) require.Equal(t, 11, got.TotalTokens) } if tc.wantReason != "" { require.Equal(t, tc.wantReason, got.Reason) } }) } } func newTianyiYunTaskRequestContext(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 }