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.
 
 
 

258 lines
8.6 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.Contains(t, models, "cdance2.0-mini-0611")
  44. require.Equal(t, "DoubaoVideoCompatibleTianyiYun", (&TaskAdaptor{}).GetChannelName())
  45. }
  46. func TestBuildRequestBodyRejectsTianyiYunDataURI(t *testing.T) {
  47. for _, mediaType := range []string{"image_url", "video_url", "audio_url"} {
  48. t.Run(mediaType, func(t *testing.T) {
  49. adaptor := &TaskAdaptor{}
  50. c := newTianyiYunTaskRequestContext(t, `{
  51. "model":"cdance2.0-0611",
  52. "metadata":{"content":[{"type":"`+mediaType+`","`+mediaType+`":{"url":"data:image/png;base64,AAAA"}}]}
  53. }`)
  54. _, err := adaptor.BuildRequestBody(c, &relaycommon.RelayInfo{ChannelMeta: &relaycommon.ChannelMeta{UpstreamModelName: "cdance2.0-0611"}})
  55. require.ErrorContains(t, err, "asset://")
  56. })
  57. }
  58. }
  59. func TestBuildRequestBodyRejectsTianyiYunUnsupportedMediaURL(t *testing.T) {
  60. for _, mediaURL := range []string{
  61. "file:///tmp/input.png",
  62. "ftp://example.test/input.png",
  63. "aGVsbG8=",
  64. } {
  65. t.Run(mediaURL, func(t *testing.T) {
  66. adaptor := &TaskAdaptor{}
  67. c := newTianyiYunTaskRequestContext(t, `{
  68. "model":"cdance2.0-0611",
  69. "metadata":{"content":[{"type":"image_url","image_url":{"url":"`+mediaURL+`"}}]}
  70. }`)
  71. _, err := adaptor.BuildRequestBody(c, &relaycommon.RelayInfo{ChannelMeta: &relaycommon.ChannelMeta{UpstreamModelName: "cdance2.0-0611"}})
  72. require.ErrorContains(t, err, "asset://")
  73. })
  74. }
  75. }
  76. func TestBuildRequestBodyPreservesTianyiYunAssetReference(t *testing.T) {
  77. adaptor := &TaskAdaptor{}
  78. c := newTianyiYunTaskRequestContext(t, `{
  79. "model":"cdance2.0-0611",
  80. "metadata":{"content":[
  81. {"type":"image_url","image_url":{"url":"asset://asset-1"},"role":"reference_image"},
  82. {"type":"video_url","video_url":{"url":"https://example.test/input.mp4"},"role":"reference_video"}
  83. ]}
  84. }`)
  85. body, err := adaptor.BuildRequestBody(c, &relaycommon.RelayInfo{ChannelMeta: &relaycommon.ChannelMeta{UpstreamModelName: "cdance2.0-0611"}})
  86. require.NoError(t, err)
  87. data, err := io.ReadAll(body)
  88. require.NoError(t, err)
  89. require.Contains(t, string(data), `"url":"asset://asset-1"`)
  90. require.Contains(t, string(data), `"url":"https://example.test/input.mp4"`)
  91. require.Less(t, strings.Index(string(data), "asset://asset-1"), strings.Index(string(data), "https://example.test/input.mp4"))
  92. }
  93. func TestBuildRequestBodyPreservesTianyiYunNativeFields(t *testing.T) {
  94. adaptor := &TaskAdaptor{}
  95. c := newTianyiYunTaskRequestContext(t, `{
  96. "model":"cdance2.0-0611",
  97. "prompt":"write a clean product video",
  98. "metadata":{
  99. "content":[
  100. {"type":"image_url","image_url":{"url":"https://example.test/input.png"},"role":"reference_image"}
  101. ],
  102. "ratio":"16:9",
  103. "duration":5,
  104. "watermark":false
  105. }
  106. }`)
  107. info := &relaycommon.RelayInfo{
  108. ChannelMeta: &relaycommon.ChannelMeta{UpstreamModelName: "cdance2.0-0611"},
  109. }
  110. body, err := adaptor.BuildRequestBody(c, info)
  111. require.NoError(t, err)
  112. data, err := io.ReadAll(body)
  113. require.NoError(t, err)
  114. require.Contains(t, string(data), `"model":"cdance2.0-0611"`)
  115. require.Contains(t, string(data), `"image_url":{"url":"https://example.test/input.png"}`)
  116. require.Contains(t, string(data), `"text":"write a clean product video"`)
  117. require.Contains(t, string(data), `"ratio":"16:9"`)
  118. require.Contains(t, string(data), `"duration":5`)
  119. require.Contains(t, string(data), `"watermark":false`)
  120. require.NotContains(t, string(data), `"seconds"`)
  121. }
  122. func TestBuildRequestBodyUsesMappedTianyiYunUpstreamModel(t *testing.T) {
  123. adaptor := &TaskAdaptor{}
  124. c := newTianyiYunTaskRequestContext(t, `{
  125. "model":"Doubao-Seedance-2.0",
  126. "prompt":"write a clean product video"
  127. }`)
  128. info := &relaycommon.RelayInfo{
  129. OriginModelName: "Doubao-Seedance-2.0",
  130. ChannelMeta: &relaycommon.ChannelMeta{
  131. UpstreamModelName: "cdance2.0-0611",
  132. IsModelMapped: true,
  133. },
  134. TaskRelayInfo: &relaycommon.TaskRelayInfo{},
  135. }
  136. body, err := adaptor.BuildRequestBody(c, info)
  137. require.NoError(t, err)
  138. data, err := io.ReadAll(body)
  139. require.NoError(t, err)
  140. require.Contains(t, string(data), `"model":"cdance2.0-0611"`)
  141. require.NotContains(t, string(data), `"model":"Doubao-Seedance-2.0"`)
  142. }
  143. func TestDoResponseRewritesIDToPublicTaskID(t *testing.T) {
  144. adaptor := &TaskAdaptor{}
  145. w := httptest.NewRecorder()
  146. c, _ := gin.CreateTestContext(w)
  147. info := &relaycommon.RelayInfo{
  148. OriginModelName: "cdance2.0-0611",
  149. TaskRelayInfo: &relaycommon.TaskRelayInfo{PublicTaskID: "task_public"},
  150. }
  151. resp := &http.Response{
  152. StatusCode: http.StatusOK,
  153. Body: io.NopCloser(strings.NewReader(`{
  154. "id":"task-upstream",
  155. "status":"queued",
  156. "created_at":1781496040,
  157. "model":"cdance2.0-0611"
  158. }`)),
  159. }
  160. upstreamID, data, taskErr := adaptor.DoResponse(c, resp, info)
  161. require.Nil(t, taskErr)
  162. require.Equal(t, "task-upstream", upstreamID)
  163. require.Contains(t, string(data), `"id":"task-upstream"`)
  164. require.Equal(t, http.StatusOK, w.Code)
  165. require.Contains(t, w.Body.String(), `"id":"task_public"`)
  166. require.Contains(t, w.Body.String(), `"status":"queued"`)
  167. require.Contains(t, w.Body.String(), `"model":"cdance2.0-0611"`)
  168. }
  169. func TestParseTaskResultMapsTianyiYunStatuses(t *testing.T) {
  170. adaptor := &TaskAdaptor{}
  171. cases := []struct {
  172. name string
  173. body string
  174. wantStatus model.TaskStatus
  175. wantURL string
  176. wantReason string
  177. wantUsage bool
  178. }{
  179. {
  180. name: "success",
  181. body: `{"id":"task-ok","status":"success","content":{"video_url":"https://example.test/out.mp4"},"usage":{"completion_tokens":7,"total_tokens":11}}`,
  182. wantStatus: model.TaskStatusSuccess,
  183. wantURL: "https://example.test/out.mp4",
  184. wantUsage: true,
  185. },
  186. {
  187. name: "succeeded",
  188. body: `{"id":"task-ok","status":"succeeded","content":{"video_url":"https://example.test/out.mp4"}}`,
  189. wantStatus: model.TaskStatusSuccess,
  190. wantURL: "https://example.test/out.mp4",
  191. },
  192. {
  193. name: "running",
  194. body: `{"id":"task-run","status":"running"}`,
  195. wantStatus: model.TaskStatusInProgress,
  196. },
  197. {
  198. name: "failed",
  199. body: `{"id":"task-fail","status":"failed","error":{"code":"BadRequest","message":"bad prompt"}}`,
  200. wantStatus: model.TaskStatusFailure,
  201. wantReason: "bad prompt",
  202. },
  203. }
  204. for _, tc := range cases {
  205. t.Run(tc.name, func(t *testing.T) {
  206. got, err := adaptor.ParseTaskResult([]byte(tc.body))
  207. require.NoError(t, err)
  208. require.Equal(t, string(tc.wantStatus), got.Status)
  209. if tc.wantURL != "" {
  210. require.Equal(t, tc.wantURL, got.Url)
  211. }
  212. if tc.wantUsage {
  213. require.Equal(t, 7, got.CompletionTokens)
  214. require.Equal(t, 11, got.TotalTokens)
  215. }
  216. if tc.wantReason != "" {
  217. require.Equal(t, tc.wantReason, got.Reason)
  218. }
  219. })
  220. }
  221. }
  222. func newTianyiYunTaskRequestContext(t *testing.T, body string) *gin.Context {
  223. t.Helper()
  224. w := httptest.NewRecorder()
  225. c, _ := gin.CreateTestContext(w)
  226. c.Request = httptest.NewRequest(http.MethodPost, "/api/v3/contents/generations/tasks", strings.NewReader(body))
  227. c.Request.Header.Set("Content-Type", "application/json")
  228. var req relaycommon.TaskSubmitReq
  229. require.NoError(t, common.Unmarshal([]byte(body), &req))
  230. relaycommon.StoreTaskRequest(c, &relaycommon.RelayInfo{}, constant.TaskActionGenerate, req)
  231. return c
  232. }