You can not select more than 25 topics Topics must start with a letter or number, can include dashes ('-') and can be up to 35 characters long.
 
 
 

207 lines
6.5 KiB

  1. package doubao_tianyiyun
  2. import (
  3. "io"
  4. "net/http"
  5. "net/http/httptest"
  6. "strings"
  7. "testing"
  8. "github.com/QuantumNous/new-api/common"
  9. "github.com/QuantumNous/new-api/constant"
  10. "github.com/QuantumNous/new-api/model"
  11. relaycommon "github.com/QuantumNous/new-api/relay/common"
  12. "github.com/gin-gonic/gin"
  13. "github.com/stretchr/testify/require"
  14. )
  15. func TestBuildRequestURLUsesTianyiYunEndpoint(t *testing.T) {
  16. adaptor := &TaskAdaptor{}
  17. adaptor.Init(&relaycommon.RelayInfo{
  18. ChannelMeta: &relaycommon.ChannelMeta{
  19. ChannelType: constant.ChannelTypeDoubaoVideoCompatibleTianyiYun,
  20. ChannelBaseUrl: "https://ai.ctaigw.cn",
  21. ApiKey: "sk-test",
  22. },
  23. })
  24. got, err := adaptor.BuildRequestURL(&relaycommon.RelayInfo{
  25. ChannelMeta: &relaycommon.ChannelMeta{ChannelBaseUrl: "https://ai.ctaigw.cn"},
  26. })
  27. require.NoError(t, err)
  28. require.Equal(t, "https://ai.ctaigw.cn/v1/contents/generations/tasks", got)
  29. }
  30. func TestFetchTaskUsesTianyiYunEndpoint(t *testing.T) {
  31. adaptor := &TaskAdaptor{}
  32. req, err := adaptor.buildFetchRequest("https://ai.ctaigw.cn", "sk-test", map[string]any{"task_id": "task-123"})
  33. require.NoError(t, err)
  34. require.Equal(t, http.MethodGet, req.Method)
  35. require.Equal(t, "https://ai.ctaigw.cn/v1/contents/generations/tasks/task-123", req.URL.String())
  36. require.Equal(t, "Bearer sk-test", req.Header.Get("Authorization"))
  37. require.Equal(t, "application/json", req.Header.Get("Accept"))
  38. }
  39. func TestGetModelListIncludesTianyiYunSeedanceModels(t *testing.T) {
  40. models := (&TaskAdaptor{}).GetModelList()
  41. require.Contains(t, models, "cdance2.0-0611")
  42. require.Contains(t, models, "cdance2.0-fast-0611")
  43. require.Equal(t, "DoubaoVideoCompatibleTianyiYun", (&TaskAdaptor{}).GetChannelName())
  44. }
  45. func TestBuildRequestBodyPreservesTianyiYunNativeFields(t *testing.T) {
  46. adaptor := &TaskAdaptor{}
  47. c := newTianyiYunTaskRequestContext(t, `{
  48. "model":"cdance2.0-0611",
  49. "prompt":"write a clean product video",
  50. "metadata":{
  51. "content":[
  52. {"type":"image_url","image_url":{"url":"https://example.test/input.png"},"role":"reference_image"}
  53. ],
  54. "ratio":"16:9",
  55. "duration":5,
  56. "watermark":false
  57. }
  58. }`)
  59. info := &relaycommon.RelayInfo{
  60. ChannelMeta: &relaycommon.ChannelMeta{UpstreamModelName: "cdance2.0-0611"},
  61. }
  62. body, err := adaptor.BuildRequestBody(c, info)
  63. require.NoError(t, err)
  64. data, err := io.ReadAll(body)
  65. require.NoError(t, err)
  66. require.Contains(t, string(data), `"model":"cdance2.0-0611"`)
  67. require.Contains(t, string(data), `"image_url":{"url":"https://example.test/input.png"}`)
  68. require.Contains(t, string(data), `"text":"write a clean product video"`)
  69. require.Contains(t, string(data), `"ratio":"16:9"`)
  70. require.Contains(t, string(data), `"duration":5`)
  71. require.Contains(t, string(data), `"watermark":false`)
  72. require.NotContains(t, string(data), `"seconds"`)
  73. }
  74. func TestBuildRequestBodyUsesMappedTianyiYunUpstreamModel(t *testing.T) {
  75. adaptor := &TaskAdaptor{}
  76. c := newTianyiYunTaskRequestContext(t, `{
  77. "model":"Doubao-Seedance-2.0",
  78. "prompt":"write a clean product video"
  79. }`)
  80. info := &relaycommon.RelayInfo{
  81. OriginModelName: "Doubao-Seedance-2.0",
  82. ChannelMeta: &relaycommon.ChannelMeta{
  83. UpstreamModelName: "cdance2.0-0611",
  84. IsModelMapped: true,
  85. },
  86. TaskRelayInfo: &relaycommon.TaskRelayInfo{},
  87. }
  88. body, err := adaptor.BuildRequestBody(c, info)
  89. require.NoError(t, err)
  90. data, err := io.ReadAll(body)
  91. require.NoError(t, err)
  92. require.Contains(t, string(data), `"model":"cdance2.0-0611"`)
  93. require.NotContains(t, string(data), `"model":"Doubao-Seedance-2.0"`)
  94. }
  95. func TestDoResponseRewritesIDToPublicTaskID(t *testing.T) {
  96. adaptor := &TaskAdaptor{}
  97. w := httptest.NewRecorder()
  98. c, _ := gin.CreateTestContext(w)
  99. info := &relaycommon.RelayInfo{
  100. OriginModelName: "cdance2.0-0611",
  101. TaskRelayInfo: &relaycommon.TaskRelayInfo{PublicTaskID: "task_public"},
  102. }
  103. resp := &http.Response{
  104. StatusCode: http.StatusOK,
  105. Body: io.NopCloser(strings.NewReader(`{
  106. "id":"task-upstream",
  107. "status":"queued",
  108. "created_at":1781496040,
  109. "model":"cdance2.0-0611"
  110. }`)),
  111. }
  112. upstreamID, data, taskErr := adaptor.DoResponse(c, resp, info)
  113. require.Nil(t, taskErr)
  114. require.Equal(t, "task-upstream", upstreamID)
  115. require.Contains(t, string(data), `"id":"task-upstream"`)
  116. require.Equal(t, http.StatusOK, w.Code)
  117. require.Contains(t, w.Body.String(), `"id":"task_public"`)
  118. require.Contains(t, w.Body.String(), `"status":"queued"`)
  119. require.Contains(t, w.Body.String(), `"model":"cdance2.0-0611"`)
  120. }
  121. func TestParseTaskResultMapsTianyiYunStatuses(t *testing.T) {
  122. adaptor := &TaskAdaptor{}
  123. cases := []struct {
  124. name string
  125. body string
  126. wantStatus model.TaskStatus
  127. wantURL string
  128. wantReason string
  129. wantUsage bool
  130. }{
  131. {
  132. name: "success",
  133. body: `{"id":"task-ok","status":"success","content":{"video_url":"https://example.test/out.mp4"},"usage":{"completion_tokens":7,"total_tokens":11}}`,
  134. wantStatus: model.TaskStatusSuccess,
  135. wantURL: "https://example.test/out.mp4",
  136. wantUsage: true,
  137. },
  138. {
  139. name: "succeeded",
  140. body: `{"id":"task-ok","status":"succeeded","content":{"video_url":"https://example.test/out.mp4"}}`,
  141. wantStatus: model.TaskStatusSuccess,
  142. wantURL: "https://example.test/out.mp4",
  143. },
  144. {
  145. name: "running",
  146. body: `{"id":"task-run","status":"running"}`,
  147. wantStatus: model.TaskStatusInProgress,
  148. },
  149. {
  150. name: "failed",
  151. body: `{"id":"task-fail","status":"failed","error":{"code":"BadRequest","message":"bad prompt"}}`,
  152. wantStatus: model.TaskStatusFailure,
  153. wantReason: "bad prompt",
  154. },
  155. }
  156. for _, tc := range cases {
  157. t.Run(tc.name, func(t *testing.T) {
  158. got, err := adaptor.ParseTaskResult([]byte(tc.body))
  159. require.NoError(t, err)
  160. require.Equal(t, string(tc.wantStatus), got.Status)
  161. if tc.wantURL != "" {
  162. require.Equal(t, tc.wantURL, got.Url)
  163. }
  164. if tc.wantUsage {
  165. require.Equal(t, 7, got.CompletionTokens)
  166. require.Equal(t, 11, got.TotalTokens)
  167. }
  168. if tc.wantReason != "" {
  169. require.Equal(t, tc.wantReason, got.Reason)
  170. }
  171. })
  172. }
  173. }
  174. func newTianyiYunTaskRequestContext(t *testing.T, body string) *gin.Context {
  175. t.Helper()
  176. w := httptest.NewRecorder()
  177. c, _ := gin.CreateTestContext(w)
  178. c.Request = httptest.NewRequest(http.MethodPost, "/api/v3/contents/generations/tasks", strings.NewReader(body))
  179. c.Request.Header.Set("Content-Type", "application/json")
  180. var req relaycommon.TaskSubmitReq
  181. require.NoError(t, common.Unmarshal([]byte(body), &req))
  182. relaycommon.StoreTaskRequest(c, &relaycommon.RelayInfo{}, constant.TaskActionGenerate, req)
  183. return c
  184. }