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 }