25개 이상의 토픽을 선택하실 수 없습니다. Topics must start with a letter or number, can include dashes ('-') and can be up to 35 characters long.
 
 
 

350 lines
12 KiB

  1. package chinamobile_seedance
  2. import (
  3. "errors"
  4. "io"
  5. "net/http"
  6. "net/http/httptest"
  7. "strings"
  8. "testing"
  9. "time"
  10. "github.com/QuantumNous/new-api/common"
  11. "github.com/QuantumNous/new-api/constant"
  12. "github.com/QuantumNous/new-api/model"
  13. relaycommon "github.com/QuantumNous/new-api/relay/common"
  14. "github.com/gin-gonic/gin"
  15. "github.com/stretchr/testify/require"
  16. )
  17. func TestRunSDKCallRecoversPanic(t *testing.T) {
  18. _, err := runSDKCall[string](time.Second, func() (string, error) {
  19. panic("attestation failed")
  20. })
  21. require.Error(t, err)
  22. require.Contains(t, err.Error(), "upstream_sdk_panic")
  23. }
  24. func TestRunSDKCallTimesOut(t *testing.T) {
  25. start := time.Now()
  26. _, err := runSDKCall[string](20*time.Millisecond, func() (string, error) {
  27. time.Sleep(200 * time.Millisecond)
  28. return "late", nil
  29. })
  30. require.Error(t, err)
  31. require.Contains(t, err.Error(), "upstream_timeout")
  32. require.Less(t, time.Since(start), 150*time.Millisecond)
  33. }
  34. func TestRunSDKCallTimeoutDoesNotPermanentlyOccupySemaphore(t *testing.T) {
  35. original := sdkSemaphore
  36. sdkSemaphore = make(chan struct{}, 1)
  37. t.Cleanup(func() {
  38. sdkSemaphore = original
  39. })
  40. block := make(chan struct{})
  41. _, err := runSDKCall[string](20*time.Millisecond, func() (string, error) {
  42. <-block
  43. return "late", nil
  44. })
  45. require.ErrorContains(t, err, "upstream_timeout")
  46. _, err = runSDKCall[string](time.Second, func() (string, error) {
  47. return "ok", nil
  48. })
  49. require.NoError(t, err)
  50. close(block)
  51. }
  52. func TestRunSDKCallReturnsError(t *testing.T) {
  53. _, err := runSDKCall[string](time.Second, func() (string, error) {
  54. return "", errors.New("upstream rejected")
  55. })
  56. require.ErrorContains(t, err, "upstream rejected")
  57. }
  58. type fakeSDK struct {
  59. createInput map[string]interface{}
  60. createID string
  61. queryTaskID string
  62. queryResult map[string]interface{}
  63. err error
  64. }
  65. func (f *fakeSDK) CreateVideoGenerationTask(data map[string]interface{}) (string, error) {
  66. f.createInput = data
  67. if f.err != nil {
  68. return "", f.err
  69. }
  70. return f.createID, nil
  71. }
  72. func (f *fakeSDK) QueryVideoGenerationTask(taskID string) (map[string]interface{}, error) {
  73. f.queryTaskID = taskID
  74. if f.err != nil {
  75. return nil, f.err
  76. }
  77. return f.queryResult, nil
  78. }
  79. func TestBuildRequestBodyNormalizesContentToInterfaceSlice(t *testing.T) {
  80. adaptor := &TaskAdaptor{}
  81. adaptor.Init(&relaycommon.RelayInfo{
  82. ChannelMeta: &relaycommon.ChannelMeta{ChannelType: constant.ChannelTypeChinaMobileSeedance},
  83. })
  84. c := newTaskRequestContext(t, `{
  85. "model":"doubao-seedance-2.0",
  86. "prompt":"current prompt",
  87. "metadata":{
  88. "content":[
  89. {"type":"video_url","video_url":{"url":"https://example.test/input.mp4"},"role":"reference_video"}
  90. ],
  91. "duration":5
  92. }
  93. }`)
  94. info := &relaycommon.RelayInfo{
  95. ChannelMeta: &relaycommon.ChannelMeta{UpstreamModelName: "doubao-seedance-2.0"},
  96. }
  97. body, err := adaptor.BuildRequestBody(c, info)
  98. require.NoError(t, err)
  99. data, err := io.ReadAll(body)
  100. require.NoError(t, err)
  101. var payload map[string]interface{}
  102. require.NoError(t, common.Unmarshal(data, &payload))
  103. content, ok := payload["content"].([]interface{})
  104. require.True(t, ok)
  105. require.Len(t, content, 2)
  106. require.Equal(t, "video_url", content[0].(map[string]interface{})["type"])
  107. require.Equal(t, "text", content[1].(map[string]interface{})["type"])
  108. require.Equal(t, "current prompt", content[1].(map[string]interface{})["text"])
  109. require.Equal(t, float64(5), payload["duration"])
  110. }
  111. func TestDoRequestCallsSDKAndReturnsSyntheticResponse(t *testing.T) {
  112. fake := &fakeSDK{createID: "upstream-task-1"}
  113. adaptor := &TaskAdaptor{
  114. newSDK: func(baseURL, apiKey, model string) (seedanceSDK, error) {
  115. require.Equal(t, "https://cm.example.com/api/v3", baseURL)
  116. require.Equal(t, "sk-test", apiKey)
  117. require.Equal(t, "doubao-seedance-2.0", model)
  118. return fake, nil
  119. },
  120. }
  121. adaptor.Init(&relaycommon.RelayInfo{
  122. ChannelMeta: &relaycommon.ChannelMeta{
  123. ChannelType: constant.ChannelTypeChinaMobileSeedance,
  124. ChannelBaseUrl: "https://cm.example.com/api/v3",
  125. ApiKey: "sk-test",
  126. },
  127. OriginModelName: "doubao-seedance-2.0",
  128. })
  129. payload := strings.NewReader(`{"model":"doubao-seedance-2.0","content":[{"type":"text","text":"hello"}]}`)
  130. resp, err := adaptor.DoRequest(&gin.Context{}, &relaycommon.RelayInfo{
  131. ChannelMeta: &relaycommon.ChannelMeta{
  132. ChannelType: constant.ChannelTypeChinaMobileSeedance,
  133. ChannelBaseUrl: "https://cm.example.com/api/v3",
  134. ApiKey: "sk-test",
  135. },
  136. OriginModelName: "doubao-seedance-2.0",
  137. }, payload)
  138. require.NoError(t, err)
  139. require.Equal(t, http.StatusOK, resp.StatusCode)
  140. data, err := io.ReadAll(resp.Body)
  141. require.NoError(t, err)
  142. require.JSONEq(t, `{"id":"upstream-task-1"}`, string(data))
  143. require.Equal(t, "doubao-seedance-2.0", fake.createInput["model"])
  144. }
  145. func TestDoRequestMapsOfficialSeedanceModelToChinaMobileDefault(t *testing.T) {
  146. fake := &fakeSDK{createID: "upstream-task-1"}
  147. adaptor := &TaskAdaptor{
  148. newSDK: func(baseURL, apiKey, model string) (seedanceSDK, error) {
  149. require.Equal(t, "doubao-seedance-2.0", model)
  150. return fake, nil
  151. },
  152. }
  153. payload := strings.NewReader(`{"model":"doubao-seedance-2-0-260128","content":[{"type":"text","text":"hello"}]}`)
  154. resp, err := adaptor.DoRequest(&gin.Context{}, &relaycommon.RelayInfo{
  155. ChannelMeta: &relaycommon.ChannelMeta{ChannelType: constant.ChannelTypeChinaMobileSeedance},
  156. OriginModelName: "doubao-seedance-2-0-260128",
  157. }, payload)
  158. require.NoError(t, err)
  159. require.Equal(t, http.StatusOK, resp.StatusCode)
  160. require.Equal(t, "doubao-seedance-2.0", fake.createInput["model"])
  161. }
  162. func TestDoRequestMapsPermissionErrorToForbiddenResponse(t *testing.T) {
  163. adaptor := &TaskAdaptor{
  164. newSDK: func(baseURL, apiKey, model string) (seedanceSDK, error) {
  165. return &fakeSDK{err: errors.New(`Failed to create video generation task: {"ErrorCode":"PERMISSION_ERROR","ErrorMessage":"Endpoint is not authorized"}`)}, nil
  166. },
  167. }
  168. payload := strings.NewReader(`{"model":"doubao-seedance-2.0","content":[{"type":"text","text":"hello"}]}`)
  169. resp, err := adaptor.DoRequest(&gin.Context{}, &relaycommon.RelayInfo{
  170. ChannelMeta: &relaycommon.ChannelMeta{ChannelType: constant.ChannelTypeChinaMobileSeedance},
  171. OriginModelName: "doubao-seedance-2.0",
  172. }, payload)
  173. require.NoError(t, err)
  174. require.Equal(t, http.StatusForbidden, resp.StatusCode)
  175. data, err := io.ReadAll(resp.Body)
  176. require.NoError(t, err)
  177. require.Contains(t, string(data), "PERMISSION_ERROR")
  178. require.Contains(t, string(data), "Endpoint is not authorized")
  179. }
  180. func TestDoRequestMapsSensitiveContentErrorToBadRequestResponse(t *testing.T) {
  181. adaptor := &TaskAdaptor{
  182. newSDK: func(baseURL, apiKey, model string) (seedanceSDK, error) {
  183. return &fakeSDK{err: errors.New(`Failed to create video generation task: {"ErrorCode":"InputVideoSensitiveContentDetected.PrivacyInformation","ErrorMessage":"input video may contain real person"}`)}, nil
  184. },
  185. }
  186. payload := strings.NewReader(`{"model":"doubao-seedance-2.0","content":[{"type":"text","text":"hello"}]}`)
  187. resp, err := adaptor.DoRequest(&gin.Context{}, &relaycommon.RelayInfo{
  188. ChannelMeta: &relaycommon.ChannelMeta{ChannelType: constant.ChannelTypeChinaMobileSeedance},
  189. OriginModelName: "doubao-seedance-2.0",
  190. }, payload)
  191. require.NoError(t, err)
  192. require.Equal(t, http.StatusBadRequest, resp.StatusCode)
  193. data, err := io.ReadAll(resp.Body)
  194. require.NoError(t, err)
  195. require.Contains(t, string(data), "InputVideoSensitiveContentDetected.PrivacyInformation")
  196. }
  197. func TestFetchTaskCallsSDKAndReturnsQueryMap(t *testing.T) {
  198. fake := &fakeSDK{queryResult: map[string]interface{}{
  199. "id": "upstream-task-1",
  200. "status": "succeeded",
  201. "content": map[string]interface{}{
  202. "video_url": "https://example.test/video.mp4",
  203. },
  204. }}
  205. adaptor := &TaskAdaptor{
  206. newSDK: func(baseURL, apiKey, model string) (seedanceSDK, error) {
  207. require.Equal(t, "https://cm.example.com/api/v3", baseURL)
  208. require.Equal(t, "sk-test", apiKey)
  209. require.Equal(t, "doubao-seedance-2.0", model)
  210. return fake, nil
  211. },
  212. }
  213. adaptor.Init(&relaycommon.RelayInfo{ChannelMeta: &relaycommon.ChannelMeta{ChannelType: constant.ChannelTypeChinaMobileSeedance}})
  214. resp, err := adaptor.FetchTask("https://cm.example.com/api/v3", "sk-test", map[string]any{
  215. "task_id": "upstream-task-1",
  216. "model": "doubao-seedance-2.0",
  217. }, "")
  218. require.NoError(t, err)
  219. require.Equal(t, "upstream-task-1", fake.queryTaskID)
  220. data, err := io.ReadAll(resp.Body)
  221. require.NoError(t, err)
  222. require.Contains(t, string(data), `"video_url":"https://example.test/video.mp4"`)
  223. }
  224. func TestFetchTaskWithoutInitUsesDefaultTimeout(t *testing.T) {
  225. fake := &fakeSDK{queryResult: map[string]interface{}{
  226. "id": "upstream-task-1",
  227. "status": "queued",
  228. }}
  229. adaptor := &TaskAdaptor{
  230. newSDK: func(baseURL, apiKey, model string) (seedanceSDK, error) {
  231. return fake, nil
  232. },
  233. }
  234. resp, err := adaptor.FetchTask("https://cm.example.com/api/v3", "sk-test", map[string]any{
  235. "task_id": "upstream-task-1",
  236. "model": "doubao-seedance-2.0",
  237. }, "")
  238. require.NoError(t, err)
  239. require.Equal(t, http.StatusOK, resp.StatusCode)
  240. }
  241. func TestDoResponseKeepsUpstreamIDWhenPublicTaskIDMissing(t *testing.T) {
  242. adaptor := &TaskAdaptor{}
  243. w := httptest.NewRecorder()
  244. c, _ := gin.CreateTestContext(w)
  245. resp := &http.Response{
  246. StatusCode: http.StatusOK,
  247. Body: io.NopCloser(strings.NewReader(`{"id":"upstream-task-1"}`)),
  248. }
  249. upstreamID, _, taskErr := adaptor.DoResponse(c, resp, &relaycommon.RelayInfo{
  250. OriginModelName: "doubao-seedance-2.0",
  251. TaskRelayInfo: &relaycommon.TaskRelayInfo{},
  252. })
  253. require.Nil(t, taskErr)
  254. require.Equal(t, "upstream-task-1", upstreamID)
  255. require.Contains(t, w.Body.String(), `"id":"upstream-task-1"`)
  256. }
  257. func TestParseTaskResultSuccessWithVideoURLAndUsage(t *testing.T) {
  258. taskInfo, err := (&TaskAdaptor{}).ParseTaskResult([]byte(`{
  259. "id":"cm-task",
  260. "status":"succeeded",
  261. "content":{"video_url":"https://example.com/video.mp4"},
  262. "usage":{"completion_tokens":12,"total_tokens":34}
  263. }`))
  264. require.NoError(t, err)
  265. require.Equal(t, model.TaskStatusSuccess, taskInfo.Status)
  266. require.Equal(t, "https://example.com/video.mp4", taskInfo.Url)
  267. require.Equal(t, 12, taskInfo.CompletionTokens)
  268. require.Equal(t, 34, taskInfo.TotalTokens)
  269. }
  270. func TestParseTaskResultSuccessWithoutUsage(t *testing.T) {
  271. taskInfo, err := (&TaskAdaptor{}).ParseTaskResult([]byte(`{
  272. "id":"cm-task",
  273. "status":"success",
  274. "content":{"video_url":"https://example.com/video.mp4"}
  275. }`))
  276. require.NoError(t, err)
  277. require.Equal(t, model.TaskStatusSuccess, taskInfo.Status)
  278. require.Equal(t, 0, taskInfo.CompletionTokens)
  279. require.Equal(t, 0, taskInfo.TotalTokens)
  280. }
  281. func TestParseTaskResultFailureFromErrorMap(t *testing.T) {
  282. taskInfo, err := (&TaskAdaptor{}).ParseTaskResult([]byte(`{
  283. "error":{"code":"bad_request","message":"invalid content"}
  284. }`))
  285. require.NoError(t, err)
  286. require.Equal(t, model.TaskStatusFailure, taskInfo.Status)
  287. require.Equal(t, "invalid content", taskInfo.Reason)
  288. }
  289. func TestParseTaskResultFailureFromErrorString(t *testing.T) {
  290. taskInfo, err := (&TaskAdaptor{}).ParseTaskResult([]byte(`{
  291. "status":"failed",
  292. "error":"quota exceeded"
  293. }`))
  294. require.NoError(t, err)
  295. require.Equal(t, model.TaskStatusFailure, taskInfo.Status)
  296. require.Equal(t, "quota exceeded", taskInfo.Reason)
  297. }
  298. func newTaskRequestContext(t *testing.T, body string) *gin.Context {
  299. t.Helper()
  300. w := httptest.NewRecorder()
  301. c, _ := gin.CreateTestContext(w)
  302. c.Request = httptest.NewRequest(http.MethodPost, "/api/v3/contents/generations/tasks", strings.NewReader(body))
  303. c.Request.Header.Set("Content-Type", "application/json")
  304. var req relaycommon.TaskSubmitReq
  305. require.NoError(t, common.Unmarshal([]byte(body), &req))
  306. relaycommon.StoreTaskRequest(c, &relaycommon.RelayInfo{}, constant.TaskActionGenerate, req)
  307. return c
  308. }