|
- 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
- }
|