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